init
This commit is contained in:
@@ -0,0 +1,9 @@
|
|||||||
|
logs
|
||||||
|
__pycache__
|
||||||
|
wd14_tagger_model
|
||||||
|
venv
|
||||||
|
*.egg-info
|
||||||
|
build
|
||||||
|
.vscode
|
||||||
|
wandb
|
||||||
|
output
|
||||||
+201
@@ -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 [2022] [kohya-ss]
|
||||||
|
|
||||||
|
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,5 @@
|
|||||||
|
# ComfyUI Flux Trainer
|
||||||
|
|
||||||
|
Currently supports LoRA training with kohya's scripts.
|
||||||
|
|
||||||
|
Original training code: https://github.com/kohya-ss/sd-scripts
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||||
|
|
||||||
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
+560
@@ -0,0 +1,560 @@
|
|||||||
|
# training with captions
|
||||||
|
# XXX dropped option: hypernetwork training
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
from multiprocessing import Value
|
||||||
|
import toml
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from library import deepspeed_utils, strategy_base
|
||||||
|
from library.device_utils import init_ipex, clean_memory_on_device
|
||||||
|
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
from accelerate.utils import set_seed
|
||||||
|
from diffusers import DDPMScheduler
|
||||||
|
|
||||||
|
from .utils import setup_logging, add_logging_arguments
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
import library.train_util as train_util
|
||||||
|
import library.config_util as config_util
|
||||||
|
from library.config_util import (
|
||||||
|
ConfigSanitizer,
|
||||||
|
BlueprintGenerator,
|
||||||
|
)
|
||||||
|
import library.custom_train_functions as custom_train_functions
|
||||||
|
from library.custom_train_functions import (
|
||||||
|
apply_snr_weight,
|
||||||
|
get_weighted_text_embeddings,
|
||||||
|
prepare_scheduler_for_custom_training,
|
||||||
|
scale_v_prediction_loss_like_noise_prediction,
|
||||||
|
apply_debiased_estimation,
|
||||||
|
)
|
||||||
|
import library.strategy_sd as strategy_sd
|
||||||
|
|
||||||
|
|
||||||
|
def train(args):
|
||||||
|
train_util.verify_training_args(args)
|
||||||
|
train_util.prepare_dataset_args(args, True)
|
||||||
|
deepspeed_utils.prepare_deepspeed_args(args)
|
||||||
|
setup_logging(args, reset=True)
|
||||||
|
|
||||||
|
cache_latents = args.cache_latents
|
||||||
|
|
||||||
|
if args.seed is not None:
|
||||||
|
set_seed(args.seed) # 乱数系列を初期化する
|
||||||
|
|
||||||
|
tokenize_strategy = strategy_sd.SdTokenizeStrategy(args.v2, args.max_token_length, args.tokenizer_cache_dir)
|
||||||
|
strategy_base.TokenizeStrategy.set_strategy(tokenize_strategy)
|
||||||
|
|
||||||
|
# prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization.
|
||||||
|
if cache_latents:
|
||||||
|
latents_caching_strategy = strategy_sd.SdSdxlLatentsCachingStrategy(
|
||||||
|
False, args.cache_latents_to_disk, args.vae_batch_size, False
|
||||||
|
)
|
||||||
|
strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy)
|
||||||
|
|
||||||
|
# データセットを準備する
|
||||||
|
if args.dataset_class is None:
|
||||||
|
blueprint_generator = BlueprintGenerator(ConfigSanitizer(False, True, False, 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:
|
||||||
|
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)
|
||||||
|
|
||||||
|
if args.debug_dataset:
|
||||||
|
train_util.debug_dataset(train_dataset_group)
|
||||||
|
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は使えません"
|
||||||
|
|
||||||
|
# acceleratorを準備する
|
||||||
|
logger.info("prepare accelerator")
|
||||||
|
accelerator = train_util.prepare_accelerator(args)
|
||||||
|
|
||||||
|
# mixed precisionに対応した型を用意しておき適宜castする
|
||||||
|
weight_dtype, save_dtype = train_util.prepare_dtype(args)
|
||||||
|
vae_dtype = torch.float32 if args.no_half_vae else weight_dtype
|
||||||
|
|
||||||
|
# モデルを読み込む
|
||||||
|
text_encoder, vae, unet, load_stable_diffusion_format = train_util.load_target_model(args, weight_dtype, accelerator)
|
||||||
|
|
||||||
|
# verify load/save model formats
|
||||||
|
if load_stable_diffusion_format:
|
||||||
|
src_stable_diffusion_ckpt = args.pretrained_model_name_or_path
|
||||||
|
src_diffusers_model_path = None
|
||||||
|
else:
|
||||||
|
src_stable_diffusion_ckpt = None
|
||||||
|
src_diffusers_model_path = args.pretrained_model_name_or_path
|
||||||
|
|
||||||
|
if args.save_model_as is None:
|
||||||
|
save_stable_diffusion_format = load_stable_diffusion_format
|
||||||
|
use_safetensors = args.use_safetensors
|
||||||
|
else:
|
||||||
|
save_stable_diffusion_format = args.save_model_as.lower() == "ckpt" or args.save_model_as.lower() == "safetensors"
|
||||||
|
use_safetensors = args.use_safetensors or ("safetensors" in args.save_model_as.lower())
|
||||||
|
|
||||||
|
# Diffusers版のxformers使用フラグを設定する関数
|
||||||
|
def set_diffusers_xformers_flag(model, valid):
|
||||||
|
# model.set_use_memory_efficient_attention_xformers(valid) # 次のリリースでなくなりそう
|
||||||
|
# pipeが自動で再帰的にset_use_memory_efficient_attention_xformersを探すんだって(;´Д`)
|
||||||
|
# U-Netだけ使う時にはどうすればいいのか……仕方ないからコピって使うか
|
||||||
|
# 0.10.2でなんか巻き戻って個別に指定するようになった(;^ω^)
|
||||||
|
|
||||||
|
# Recursively walk through all the children.
|
||||||
|
# Any children which exposes the set_use_memory_efficient_attention_xformers method
|
||||||
|
# gets the message
|
||||||
|
def fn_recursive_set_mem_eff(module: torch.nn.Module):
|
||||||
|
if hasattr(module, "set_use_memory_efficient_attention_xformers"):
|
||||||
|
module.set_use_memory_efficient_attention_xformers(valid)
|
||||||
|
|
||||||
|
for child in module.children():
|
||||||
|
fn_recursive_set_mem_eff(child)
|
||||||
|
|
||||||
|
fn_recursive_set_mem_eff(model)
|
||||||
|
|
||||||
|
# モデルに xformers とか memory efficient attention を組み込む
|
||||||
|
if args.diffusers_xformers:
|
||||||
|
accelerator.print("Use xformers by Diffusers")
|
||||||
|
set_diffusers_xformers_flag(unet, True)
|
||||||
|
else:
|
||||||
|
# Windows版のxformersはfloatで学習できないのでxformersを使わない設定も可能にしておく必要がある
|
||||||
|
accelerator.print("Disable Diffusers' xformers")
|
||||||
|
set_diffusers_xformers_flag(unet, False)
|
||||||
|
train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers, args.sdpa)
|
||||||
|
|
||||||
|
# 学習を準備する
|
||||||
|
if cache_latents:
|
||||||
|
vae.to(accelerator.device, dtype=vae_dtype)
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
vae.eval()
|
||||||
|
|
||||||
|
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process)
|
||||||
|
|
||||||
|
vae.to("cpu")
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
|
||||||
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
|
# 学習を準備する:モデルを適切な状態にする
|
||||||
|
training_models = []
|
||||||
|
if args.gradient_checkpointing:
|
||||||
|
unet.enable_gradient_checkpointing()
|
||||||
|
training_models.append(unet)
|
||||||
|
|
||||||
|
if args.train_text_encoder:
|
||||||
|
accelerator.print("enable text encoder training")
|
||||||
|
if args.gradient_checkpointing:
|
||||||
|
text_encoder.gradient_checkpointing_enable()
|
||||||
|
training_models.append(text_encoder)
|
||||||
|
else:
|
||||||
|
text_encoder.to(accelerator.device, dtype=weight_dtype)
|
||||||
|
text_encoder.requires_grad_(False) # text encoderは学習しない
|
||||||
|
if args.gradient_checkpointing:
|
||||||
|
text_encoder.gradient_checkpointing_enable()
|
||||||
|
text_encoder.train() # required for gradient_checkpointing
|
||||||
|
else:
|
||||||
|
text_encoder.eval()
|
||||||
|
|
||||||
|
text_encoding_strategy = strategy_sd.SdTextEncodingStrategy(args.clip_skip)
|
||||||
|
strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy)
|
||||||
|
|
||||||
|
if not cache_latents:
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
vae.eval()
|
||||||
|
vae.to(accelerator.device, dtype=vae_dtype)
|
||||||
|
|
||||||
|
for m in training_models:
|
||||||
|
m.requires_grad_(True)
|
||||||
|
|
||||||
|
trainable_params = []
|
||||||
|
if args.learning_rate_te is None or not args.train_text_encoder:
|
||||||
|
for m in training_models:
|
||||||
|
trainable_params.extend(m.parameters())
|
||||||
|
else:
|
||||||
|
trainable_params = [
|
||||||
|
{"params": list(unet.parameters()), "lr": args.learning_rate},
|
||||||
|
{"params": list(text_encoder.parameters()), "lr": args.learning_rate_te},
|
||||||
|
]
|
||||||
|
|
||||||
|
# 学習に必要なクラスを準備する
|
||||||
|
accelerator.print("prepare optimizer, data loader etc.")
|
||||||
|
_, _, optimizer = train_util.get_optimizer(args, trainable_params=trainable_params)
|
||||||
|
|
||||||
|
# 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を用意する
|
||||||
|
lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)
|
||||||
|
|
||||||
|
# 実験的機能:勾配も含めたfp16学習を行う モデル全体をfp16にする
|
||||||
|
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.")
|
||||||
|
unet.to(weight_dtype)
|
||||||
|
text_encoder.to(weight_dtype)
|
||||||
|
|
||||||
|
if args.deepspeed:
|
||||||
|
if args.train_text_encoder:
|
||||||
|
ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet, text_encoder=text_encoder)
|
||||||
|
else:
|
||||||
|
ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet)
|
||||||
|
ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||||
|
ds_model, optimizer, train_dataloader, lr_scheduler
|
||||||
|
)
|
||||||
|
training_models = [ds_model]
|
||||||
|
else:
|
||||||
|
# acceleratorがなんかよろしくやってくれるらしい
|
||||||
|
if args.train_text_encoder:
|
||||||
|
unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||||
|
unet, text_encoder, optimizer, train_dataloader, lr_scheduler
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(unet, optimizer, train_dataloader, lr_scheduler)
|
||||||
|
|
||||||
|
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
|
||||||
|
if args.full_fp16:
|
||||||
|
train_util.patch_accelerator_for_fp16_training(accelerator)
|
||||||
|
|
||||||
|
# resumeする
|
||||||
|
train_util.resume_from_local_or_hf_if_specified(accelerator, args)
|
||||||
|
|
||||||
|
# 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 / バッチサイズ: {args.train_batch_size}")
|
||||||
|
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 = DDPMScheduler(
|
||||||
|
beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False
|
||||||
|
)
|
||||||
|
prepare_scheduler_for_custom_training(noise_scheduler, accelerator.device)
|
||||||
|
if args.zero_terminal_snr:
|
||||||
|
custom_train_functions.fix_noise_scheduler_betas_for_zero_terminal_snr(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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# For --sample_at_first
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, 0, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
loss_recorder = train_util.LossRecorder()
|
||||||
|
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
|
||||||
|
with accelerator.accumulate(*training_models):
|
||||||
|
with torch.no_grad():
|
||||||
|
if "latents" in batch and batch["latents"] is not None:
|
||||||
|
latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)
|
||||||
|
else:
|
||||||
|
# latentに変換
|
||||||
|
latents = vae.encode(batch["images"].to(dtype=vae_dtype)).latent_dist.sample().to(weight_dtype)
|
||||||
|
latents = latents * 0.18215
|
||||||
|
b_size = latents.shape[0]
|
||||||
|
|
||||||
|
with torch.set_grad_enabled(args.train_text_encoder):
|
||||||
|
# Get the text embedding for conditioning
|
||||||
|
if args.weighted_captions:
|
||||||
|
# TODO move to strategy_sd.py
|
||||||
|
encoder_hidden_states = get_weighted_text_embeddings(
|
||||||
|
tokenize_strategy.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_ids = batch["input_ids_list"][0].to(accelerator.device)
|
||||||
|
encoder_hidden_states = text_encoding_strategy.encode_tokens(
|
||||||
|
tokenize_strategy, [text_encoder], [input_ids]
|
||||||
|
)[0]
|
||||||
|
if args.full_fp16:
|
||||||
|
encoder_hidden_states = encoder_hidden_states.to(weight_dtype)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
)
|
||||||
|
|
||||||
|
# Predict the noise residual
|
||||||
|
with accelerator.autocast():
|
||||||
|
noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample
|
||||||
|
|
||||||
|
if args.v_parameterization:
|
||||||
|
# v-parameterization training
|
||||||
|
target = noise_scheduler.get_velocity(latents, noise, timesteps)
|
||||||
|
else:
|
||||||
|
target = noise
|
||||||
|
|
||||||
|
if args.min_snr_gamma or args.scale_v_pred_loss_like_noise_pred or args.debiased_estimation_loss:
|
||||||
|
# do not mean over batch dimension for snr weight or scale v-pred loss
|
||||||
|
loss = train_util.conditional_loss(
|
||||||
|
noise_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=huber_c
|
||||||
|
)
|
||||||
|
loss = loss.mean([1, 2, 3])
|
||||||
|
|
||||||
|
if args.min_snr_gamma:
|
||||||
|
loss = apply_snr_weight(loss, timesteps, noise_scheduler, args.min_snr_gamma, args.v_parameterization)
|
||||||
|
if args.scale_v_pred_loss_like_noise_pred:
|
||||||
|
loss = scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler)
|
||||||
|
if args.debiased_estimation_loss:
|
||||||
|
loss = apply_debiased_estimation(loss, timesteps, noise_scheduler)
|
||||||
|
|
||||||
|
loss = loss.mean() # mean over batch dimension
|
||||||
|
else:
|
||||||
|
loss = train_util.conditional_loss(
|
||||||
|
noise_pred.float(), target.float(), reduction="mean", loss_type=args.loss_type, huber_c=huber_c
|
||||||
|
)
|
||||||
|
|
||||||
|
accelerator.backward(loss)
|
||||||
|
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)
|
||||||
|
|
||||||
|
# Checks if the accelerator has performed an optimization step behind the scenes
|
||||||
|
if accelerator.sync_gradients:
|
||||||
|
progress_bar.update(1)
|
||||||
|
global_step += 1
|
||||||
|
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, None, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
# 指定ステップごとにモデルを保存
|
||||||
|
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:
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||||
|
args,
|
||||||
|
False,
|
||||||
|
accelerator,
|
||||||
|
src_path,
|
||||||
|
save_stable_diffusion_format,
|
||||||
|
use_safetensors,
|
||||||
|
save_dtype,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
accelerator.unwrap_model(text_encoder),
|
||||||
|
accelerator.unwrap_model(unet),
|
||||||
|
vae,
|
||||||
|
)
|
||||||
|
|
||||||
|
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:
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||||
|
args,
|
||||||
|
True,
|
||||||
|
accelerator,
|
||||||
|
src_path,
|
||||||
|
save_stable_diffusion_format,
|
||||||
|
use_safetensors,
|
||||||
|
save_dtype,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
accelerator.unwrap_model(text_encoder),
|
||||||
|
accelerator.unwrap_model(unet),
|
||||||
|
vae,
|
||||||
|
)
|
||||||
|
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, epoch + 1, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
is_main_process = accelerator.is_main_process
|
||||||
|
if is_main_process:
|
||||||
|
unet = accelerator.unwrap_model(unet)
|
||||||
|
text_encoder = accelerator.unwrap_model(text_encoder)
|
||||||
|
|
||||||
|
accelerator.end_training()
|
||||||
|
|
||||||
|
if is_main_process and (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:
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_train_end(
|
||||||
|
args, src_path, save_stable_diffusion_format, use_safetensors, save_dtype, epoch, global_step, text_encoder, unet, vae
|
||||||
|
)
|
||||||
|
logger.info("model saved.")
|
||||||
|
|
||||||
|
|
||||||
|
def setup_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
add_logging_arguments(parser)
|
||||||
|
train_util.add_sd_models_arguments(parser)
|
||||||
|
train_util.add_dataset_arguments(parser, False, True, True)
|
||||||
|
train_util.add_training_arguments(parser, False)
|
||||||
|
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)
|
||||||
|
custom_train_functions.add_custom_train_arguments(parser)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--diffusers_xformers", action="store_true", help="use xformers by diffusers / Diffusersでxformersを使用する"
|
||||||
|
)
|
||||||
|
parser.add_argument("--train_text_encoder", action="store_true", help="train text encoder / text encoderも学習する")
|
||||||
|
parser.add_argument(
|
||||||
|
"--learning_rate_te",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="learning rate for text encoder, default is same as unet / Text Encoderの学習率、デフォルトはunetと同じ",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--no_half_vae",
|
||||||
|
action="store_true",
|
||||||
|
help="do not use fp16/bf16 VAE in mixed precision (use float VAE) / mixed precisionでも fp16/bf16 VAEを使わずfloat VAEを使う",
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
import math
|
||||||
|
import torch
|
||||||
|
from transformers import Adafactor
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def adafactor_step_param(self, p, group):
|
||||||
|
if p.grad is None:
|
||||||
|
return
|
||||||
|
grad = p.grad
|
||||||
|
if grad.dtype in {torch.float16, torch.bfloat16}:
|
||||||
|
grad = grad.float()
|
||||||
|
if grad.is_sparse:
|
||||||
|
raise RuntimeError("Adafactor does not support sparse gradients.")
|
||||||
|
|
||||||
|
state = self.state[p]
|
||||||
|
grad_shape = grad.shape
|
||||||
|
|
||||||
|
factored, use_first_moment = Adafactor._get_options(group, grad_shape)
|
||||||
|
# State Initialization
|
||||||
|
if len(state) == 0:
|
||||||
|
state["step"] = 0
|
||||||
|
|
||||||
|
if use_first_moment:
|
||||||
|
# Exponential moving average of gradient values
|
||||||
|
state["exp_avg"] = torch.zeros_like(grad)
|
||||||
|
if factored:
|
||||||
|
state["exp_avg_sq_row"] = torch.zeros(grad_shape[:-1]).to(grad)
|
||||||
|
state["exp_avg_sq_col"] = torch.zeros(grad_shape[:-2] + grad_shape[-1:]).to(grad)
|
||||||
|
else:
|
||||||
|
state["exp_avg_sq"] = torch.zeros_like(grad)
|
||||||
|
|
||||||
|
state["RMS"] = 0
|
||||||
|
else:
|
||||||
|
if use_first_moment:
|
||||||
|
state["exp_avg"] = state["exp_avg"].to(grad)
|
||||||
|
if factored:
|
||||||
|
state["exp_avg_sq_row"] = state["exp_avg_sq_row"].to(grad)
|
||||||
|
state["exp_avg_sq_col"] = state["exp_avg_sq_col"].to(grad)
|
||||||
|
else:
|
||||||
|
state["exp_avg_sq"] = state["exp_avg_sq"].to(grad)
|
||||||
|
|
||||||
|
p_data_fp32 = p
|
||||||
|
if p.dtype in {torch.float16, torch.bfloat16}:
|
||||||
|
p_data_fp32 = p_data_fp32.float()
|
||||||
|
|
||||||
|
state["step"] += 1
|
||||||
|
state["RMS"] = Adafactor._rms(p_data_fp32)
|
||||||
|
lr = Adafactor._get_lr(group, state)
|
||||||
|
|
||||||
|
beta2t = 1.0 - math.pow(state["step"], group["decay_rate"])
|
||||||
|
update = (grad ** 2) + group["eps"][0]
|
||||||
|
if factored:
|
||||||
|
exp_avg_sq_row = state["exp_avg_sq_row"]
|
||||||
|
exp_avg_sq_col = state["exp_avg_sq_col"]
|
||||||
|
|
||||||
|
exp_avg_sq_row.mul_(beta2t).add_(update.mean(dim=-1), alpha=(1.0 - beta2t))
|
||||||
|
exp_avg_sq_col.mul_(beta2t).add_(update.mean(dim=-2), alpha=(1.0 - beta2t))
|
||||||
|
|
||||||
|
# Approximation of exponential moving average of square of gradient
|
||||||
|
update = Adafactor._approx_sq_grad(exp_avg_sq_row, exp_avg_sq_col)
|
||||||
|
update.mul_(grad)
|
||||||
|
else:
|
||||||
|
exp_avg_sq = state["exp_avg_sq"]
|
||||||
|
|
||||||
|
exp_avg_sq.mul_(beta2t).add_(update, alpha=(1.0 - beta2t))
|
||||||
|
update = exp_avg_sq.rsqrt().mul_(grad)
|
||||||
|
|
||||||
|
update.div_((Adafactor._rms(update) / group["clip_threshold"]).clamp_(min=1.0))
|
||||||
|
update.mul_(lr)
|
||||||
|
|
||||||
|
if use_first_moment:
|
||||||
|
exp_avg = state["exp_avg"]
|
||||||
|
exp_avg.mul_(group["beta1"]).add_(update, alpha=(1 - group["beta1"]))
|
||||||
|
update = exp_avg
|
||||||
|
|
||||||
|
if group["weight_decay"] != 0:
|
||||||
|
p_data_fp32.add_(p_data_fp32, alpha=(-group["weight_decay"] * lr))
|
||||||
|
|
||||||
|
p_data_fp32.add_(-update)
|
||||||
|
|
||||||
|
if p.dtype in {torch.float16, torch.bfloat16}:
|
||||||
|
p.copy_(p_data_fp32)
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def adafactor_step(self, closure=None):
|
||||||
|
"""
|
||||||
|
Performs a single optimization step
|
||||||
|
|
||||||
|
Arguments:
|
||||||
|
closure (callable, optional): A closure that reevaluates the model
|
||||||
|
and returns the loss.
|
||||||
|
"""
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
for p in group["params"]:
|
||||||
|
adafactor_step_param(self, p, group)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
def patch_adafactor_fused(optimizer: Adafactor):
|
||||||
|
optimizer.step_param = adafactor_step_param.__get__(optimizer)
|
||||||
|
optimizer.step = adafactor_step.__get__(optimizer)
|
||||||
@@ -0,0 +1,227 @@
|
|||||||
|
import math
|
||||||
|
from typing import Any
|
||||||
|
from einops import rearrange
|
||||||
|
import torch
|
||||||
|
from diffusers.models.attention_processor import Attention
|
||||||
|
|
||||||
|
|
||||||
|
# flash attention forwards and backwards
|
||||||
|
|
||||||
|
# https://arxiv.org/abs/2205.14135
|
||||||
|
|
||||||
|
EPSILON = 1e-6
|
||||||
|
|
||||||
|
|
||||||
|
class FlashAttentionFunction(torch.autograd.function.Function):
|
||||||
|
@staticmethod
|
||||||
|
@torch.no_grad()
|
||||||
|
def forward(ctx, q, k, v, mask, causal, q_bucket_size, k_bucket_size):
|
||||||
|
"""Algorithm 2 in the paper"""
|
||||||
|
|
||||||
|
device = q.device
|
||||||
|
dtype = q.dtype
|
||||||
|
max_neg_value = -torch.finfo(q.dtype).max
|
||||||
|
qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)
|
||||||
|
|
||||||
|
o = torch.zeros_like(q)
|
||||||
|
all_row_sums = torch.zeros((*q.shape[:-1], 1), dtype=dtype, device=device)
|
||||||
|
all_row_maxes = torch.full(
|
||||||
|
(*q.shape[:-1], 1), max_neg_value, dtype=dtype, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
scale = q.shape[-1] ** -0.5
|
||||||
|
|
||||||
|
if mask is None:
|
||||||
|
mask = (None,) * math.ceil(q.shape[-2] / q_bucket_size)
|
||||||
|
else:
|
||||||
|
mask = rearrange(mask, "b n -> b 1 1 n")
|
||||||
|
mask = mask.split(q_bucket_size, dim=-1)
|
||||||
|
|
||||||
|
row_splits = zip(
|
||||||
|
q.split(q_bucket_size, dim=-2),
|
||||||
|
o.split(q_bucket_size, dim=-2),
|
||||||
|
mask,
|
||||||
|
all_row_sums.split(q_bucket_size, dim=-2),
|
||||||
|
all_row_maxes.split(q_bucket_size, dim=-2),
|
||||||
|
)
|
||||||
|
|
||||||
|
for ind, (qc, oc, row_mask, row_sums, row_maxes) in enumerate(row_splits):
|
||||||
|
q_start_index = ind * q_bucket_size - qk_len_diff
|
||||||
|
|
||||||
|
col_splits = zip(
|
||||||
|
k.split(k_bucket_size, dim=-2),
|
||||||
|
v.split(k_bucket_size, dim=-2),
|
||||||
|
)
|
||||||
|
|
||||||
|
for k_ind, (kc, vc) in enumerate(col_splits):
|
||||||
|
k_start_index = k_ind * k_bucket_size
|
||||||
|
|
||||||
|
attn_weights = (
|
||||||
|
torch.einsum("... i d, ... j d -> ... i j", qc, kc) * scale
|
||||||
|
)
|
||||||
|
|
||||||
|
if row_mask is not None:
|
||||||
|
attn_weights.masked_fill_(~row_mask, max_neg_value)
|
||||||
|
|
||||||
|
if causal and q_start_index < (k_start_index + k_bucket_size - 1):
|
||||||
|
causal_mask = torch.ones(
|
||||||
|
(qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device
|
||||||
|
).triu(q_start_index - k_start_index + 1)
|
||||||
|
attn_weights.masked_fill_(causal_mask, max_neg_value)
|
||||||
|
|
||||||
|
block_row_maxes = attn_weights.amax(dim=-1, keepdims=True)
|
||||||
|
attn_weights -= block_row_maxes
|
||||||
|
exp_weights = torch.exp(attn_weights)
|
||||||
|
|
||||||
|
if row_mask is not None:
|
||||||
|
exp_weights.masked_fill_(~row_mask, 0.0)
|
||||||
|
|
||||||
|
block_row_sums = exp_weights.sum(dim=-1, keepdims=True).clamp(
|
||||||
|
min=EPSILON
|
||||||
|
)
|
||||||
|
|
||||||
|
new_row_maxes = torch.maximum(block_row_maxes, row_maxes)
|
||||||
|
|
||||||
|
exp_values = torch.einsum(
|
||||||
|
"... i j, ... j d -> ... i d", exp_weights, vc
|
||||||
|
)
|
||||||
|
|
||||||
|
exp_row_max_diff = torch.exp(row_maxes - new_row_maxes)
|
||||||
|
exp_block_row_max_diff = torch.exp(block_row_maxes - new_row_maxes)
|
||||||
|
|
||||||
|
new_row_sums = (
|
||||||
|
exp_row_max_diff * row_sums
|
||||||
|
+ exp_block_row_max_diff * block_row_sums
|
||||||
|
)
|
||||||
|
|
||||||
|
oc.mul_((row_sums / new_row_sums) * exp_row_max_diff).add_(
|
||||||
|
(exp_block_row_max_diff / new_row_sums) * exp_values
|
||||||
|
)
|
||||||
|
|
||||||
|
row_maxes.copy_(new_row_maxes)
|
||||||
|
row_sums.copy_(new_row_sums)
|
||||||
|
|
||||||
|
ctx.args = (causal, scale, mask, q_bucket_size, k_bucket_size)
|
||||||
|
ctx.save_for_backward(q, k, v, o, all_row_sums, all_row_maxes)
|
||||||
|
|
||||||
|
return o
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@torch.no_grad()
|
||||||
|
def backward(ctx, do):
|
||||||
|
"""Algorithm 4 in the paper"""
|
||||||
|
|
||||||
|
causal, scale, mask, q_bucket_size, k_bucket_size = ctx.args
|
||||||
|
q, k, v, o, l, m = ctx.saved_tensors
|
||||||
|
|
||||||
|
device = q.device
|
||||||
|
|
||||||
|
max_neg_value = -torch.finfo(q.dtype).max
|
||||||
|
qk_len_diff = max(k.shape[-2] - q.shape[-2], 0)
|
||||||
|
|
||||||
|
dq = torch.zeros_like(q)
|
||||||
|
dk = torch.zeros_like(k)
|
||||||
|
dv = torch.zeros_like(v)
|
||||||
|
|
||||||
|
row_splits = zip(
|
||||||
|
q.split(q_bucket_size, dim=-2),
|
||||||
|
o.split(q_bucket_size, dim=-2),
|
||||||
|
do.split(q_bucket_size, dim=-2),
|
||||||
|
mask,
|
||||||
|
l.split(q_bucket_size, dim=-2),
|
||||||
|
m.split(q_bucket_size, dim=-2),
|
||||||
|
dq.split(q_bucket_size, dim=-2),
|
||||||
|
)
|
||||||
|
|
||||||
|
for ind, (qc, oc, doc, row_mask, lc, mc, dqc) in enumerate(row_splits):
|
||||||
|
q_start_index = ind * q_bucket_size - qk_len_diff
|
||||||
|
|
||||||
|
col_splits = zip(
|
||||||
|
k.split(k_bucket_size, dim=-2),
|
||||||
|
v.split(k_bucket_size, dim=-2),
|
||||||
|
dk.split(k_bucket_size, dim=-2),
|
||||||
|
dv.split(k_bucket_size, dim=-2),
|
||||||
|
)
|
||||||
|
|
||||||
|
for k_ind, (kc, vc, dkc, dvc) in enumerate(col_splits):
|
||||||
|
k_start_index = k_ind * k_bucket_size
|
||||||
|
|
||||||
|
attn_weights = (
|
||||||
|
torch.einsum("... i d, ... j d -> ... i j", qc, kc) * scale
|
||||||
|
)
|
||||||
|
|
||||||
|
if causal and q_start_index < (k_start_index + k_bucket_size - 1):
|
||||||
|
causal_mask = torch.ones(
|
||||||
|
(qc.shape[-2], kc.shape[-2]), dtype=torch.bool, device=device
|
||||||
|
).triu(q_start_index - k_start_index + 1)
|
||||||
|
attn_weights.masked_fill_(causal_mask, max_neg_value)
|
||||||
|
|
||||||
|
exp_attn_weights = torch.exp(attn_weights - mc)
|
||||||
|
|
||||||
|
if row_mask is not None:
|
||||||
|
exp_attn_weights.masked_fill_(~row_mask, 0.0)
|
||||||
|
|
||||||
|
p = exp_attn_weights / lc
|
||||||
|
|
||||||
|
dv_chunk = torch.einsum("... i j, ... i d -> ... j d", p, doc)
|
||||||
|
dp = torch.einsum("... i d, ... j d -> ... i j", doc, vc)
|
||||||
|
|
||||||
|
D = (doc * oc).sum(dim=-1, keepdims=True)
|
||||||
|
ds = p * scale * (dp - D)
|
||||||
|
|
||||||
|
dq_chunk = torch.einsum("... i j, ... j d -> ... i d", ds, kc)
|
||||||
|
dk_chunk = torch.einsum("... i j, ... i d -> ... j d", ds, qc)
|
||||||
|
|
||||||
|
dqc.add_(dq_chunk)
|
||||||
|
dkc.add_(dk_chunk)
|
||||||
|
dvc.add_(dv_chunk)
|
||||||
|
|
||||||
|
return dq, dk, dv, None, None, None, None
|
||||||
|
|
||||||
|
|
||||||
|
class FlashAttnProcessor:
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
attn: Attention,
|
||||||
|
hidden_states,
|
||||||
|
encoder_hidden_states=None,
|
||||||
|
attention_mask=None,
|
||||||
|
) -> Any:
|
||||||
|
q_bucket_size = 512
|
||||||
|
k_bucket_size = 1024
|
||||||
|
|
||||||
|
h = attn.heads
|
||||||
|
q = attn.to_q(hidden_states)
|
||||||
|
|
||||||
|
encoder_hidden_states = (
|
||||||
|
encoder_hidden_states
|
||||||
|
if encoder_hidden_states is not None
|
||||||
|
else hidden_states
|
||||||
|
)
|
||||||
|
encoder_hidden_states = encoder_hidden_states.to(hidden_states.dtype)
|
||||||
|
|
||||||
|
if hasattr(attn, "hypernetwork") and attn.hypernetwork is not None:
|
||||||
|
context_k, context_v = attn.hypernetwork.forward(
|
||||||
|
hidden_states, encoder_hidden_states
|
||||||
|
)
|
||||||
|
context_k = context_k.to(hidden_states.dtype)
|
||||||
|
context_v = context_v.to(hidden_states.dtype)
|
||||||
|
else:
|
||||||
|
context_k = encoder_hidden_states
|
||||||
|
context_v = encoder_hidden_states
|
||||||
|
|
||||||
|
k = attn.to_k(context_k)
|
||||||
|
v = attn.to_v(context_v)
|
||||||
|
del encoder_hidden_states, hidden_states
|
||||||
|
|
||||||
|
q, k, v = map(lambda t: rearrange(t, "b n (h d) -> b h n d", h=h), (q, k, v))
|
||||||
|
|
||||||
|
out = FlashAttentionFunction.apply(
|
||||||
|
q, k, v, attention_mask, False, q_bucket_size, k_bucket_size
|
||||||
|
)
|
||||||
|
|
||||||
|
out = rearrange(out, "b h n d -> b n (h d)")
|
||||||
|
|
||||||
|
out = attn.to_out[0](out)
|
||||||
|
out = attn.to_out[1](out)
|
||||||
|
return out
|
||||||
@@ -0,0 +1,720 @@
|
|||||||
|
import argparse
|
||||||
|
from dataclasses import (
|
||||||
|
asdict,
|
||||||
|
dataclass,
|
||||||
|
)
|
||||||
|
import functools
|
||||||
|
import random
|
||||||
|
from textwrap import dedent, indent
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
# from toolz import curry
|
||||||
|
from typing import (
|
||||||
|
List,
|
||||||
|
Optional,
|
||||||
|
Sequence,
|
||||||
|
Tuple,
|
||||||
|
Union,
|
||||||
|
)
|
||||||
|
|
||||||
|
import toml
|
||||||
|
import voluptuous
|
||||||
|
from voluptuous import (
|
||||||
|
Any,
|
||||||
|
ExactSequence,
|
||||||
|
MultipleInvalid,
|
||||||
|
Object,
|
||||||
|
Required,
|
||||||
|
Schema,
|
||||||
|
)
|
||||||
|
from transformers import CLIPTokenizer
|
||||||
|
|
||||||
|
from . import train_util
|
||||||
|
from .train_util import (
|
||||||
|
DreamBoothSubset,
|
||||||
|
FineTuningSubset,
|
||||||
|
ControlNetSubset,
|
||||||
|
DreamBoothDataset,
|
||||||
|
FineTuningDataset,
|
||||||
|
ControlNetDataset,
|
||||||
|
DatasetGroup,
|
||||||
|
)
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def add_config_arguments(parser: argparse.ArgumentParser):
|
||||||
|
parser.add_argument(
|
||||||
|
"--dataset_config", type=Path, default=None, help="config file for detail settings / 詳細な設定用の設定ファイル"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: inherit Params class in Subset, Dataset
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BaseSubsetParams:
|
||||||
|
image_dir: Optional[str] = None
|
||||||
|
num_repeats: int = 1
|
||||||
|
shuffle_caption: bool = False
|
||||||
|
caption_separator: str = (",",)
|
||||||
|
keep_tokens: int = 0
|
||||||
|
keep_tokens_separator: str = (None,)
|
||||||
|
secondary_separator: Optional[str] = None
|
||||||
|
enable_wildcard: bool = False
|
||||||
|
color_aug: bool = False
|
||||||
|
flip_aug: bool = False
|
||||||
|
face_crop_aug_range: Optional[Tuple[float, float]] = None
|
||||||
|
random_crop: bool = False
|
||||||
|
caption_prefix: Optional[str] = None
|
||||||
|
caption_suffix: Optional[str] = None
|
||||||
|
caption_dropout_rate: float = 0.0
|
||||||
|
caption_dropout_every_n_epochs: int = 0
|
||||||
|
caption_tag_dropout_rate: float = 0.0
|
||||||
|
token_warmup_min: int = 1
|
||||||
|
token_warmup_step: float = 0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DreamBoothSubsetParams(BaseSubsetParams):
|
||||||
|
is_reg: bool = False
|
||||||
|
class_tokens: Optional[str] = None
|
||||||
|
caption_extension: str = ".caption"
|
||||||
|
cache_info: bool = False
|
||||||
|
alpha_mask: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FineTuningSubsetParams(BaseSubsetParams):
|
||||||
|
metadata_file: Optional[str] = None
|
||||||
|
alpha_mask: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ControlNetSubsetParams(BaseSubsetParams):
|
||||||
|
conditioning_data_dir: str = None
|
||||||
|
caption_extension: str = ".caption"
|
||||||
|
cache_info: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BaseDatasetParams:
|
||||||
|
resolution: Optional[Tuple[int, int]] = None
|
||||||
|
network_multiplier: float = 1.0
|
||||||
|
debug_dataset: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DreamBoothDatasetParams(BaseDatasetParams):
|
||||||
|
batch_size: int = 1
|
||||||
|
enable_bucket: bool = False
|
||||||
|
min_bucket_reso: int = 256
|
||||||
|
max_bucket_reso: int = 1024
|
||||||
|
bucket_reso_steps: int = 64
|
||||||
|
bucket_no_upscale: bool = False
|
||||||
|
prior_loss_weight: float = 1.0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FineTuningDatasetParams(BaseDatasetParams):
|
||||||
|
batch_size: int = 1
|
||||||
|
enable_bucket: bool = False
|
||||||
|
min_bucket_reso: int = 256
|
||||||
|
max_bucket_reso: int = 1024
|
||||||
|
bucket_reso_steps: int = 64
|
||||||
|
bucket_no_upscale: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ControlNetDatasetParams(BaseDatasetParams):
|
||||||
|
batch_size: int = 1
|
||||||
|
enable_bucket: bool = False
|
||||||
|
min_bucket_reso: int = 256
|
||||||
|
max_bucket_reso: int = 1024
|
||||||
|
bucket_reso_steps: int = 64
|
||||||
|
bucket_no_upscale: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SubsetBlueprint:
|
||||||
|
params: Union[DreamBoothSubsetParams, FineTuningSubsetParams]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DatasetBlueprint:
|
||||||
|
is_dreambooth: bool
|
||||||
|
is_controlnet: bool
|
||||||
|
params: Union[DreamBoothDatasetParams, FineTuningDatasetParams]
|
||||||
|
subsets: Sequence[SubsetBlueprint]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DatasetGroupBlueprint:
|
||||||
|
datasets: Sequence[DatasetBlueprint]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Blueprint:
|
||||||
|
dataset_group: DatasetGroupBlueprint
|
||||||
|
|
||||||
|
|
||||||
|
class ConfigSanitizer:
|
||||||
|
# @curry
|
||||||
|
@staticmethod
|
||||||
|
def __validate_and_convert_twodim(klass, value: Sequence) -> Tuple:
|
||||||
|
Schema(ExactSequence([klass, klass]))(value)
|
||||||
|
return tuple(value)
|
||||||
|
|
||||||
|
# @curry
|
||||||
|
@staticmethod
|
||||||
|
def __validate_and_convert_scalar_or_twodim(klass, value: Union[float, Sequence]) -> Tuple:
|
||||||
|
Schema(Any(klass, ExactSequence([klass, klass])))(value)
|
||||||
|
try:
|
||||||
|
Schema(klass)(value)
|
||||||
|
return (value, value)
|
||||||
|
except:
|
||||||
|
return ConfigSanitizer.__validate_and_convert_twodim(klass, value)
|
||||||
|
|
||||||
|
# subset schema
|
||||||
|
SUBSET_ASCENDABLE_SCHEMA = {
|
||||||
|
"color_aug": bool,
|
||||||
|
"face_crop_aug_range": functools.partial(__validate_and_convert_twodim.__func__, float),
|
||||||
|
"flip_aug": bool,
|
||||||
|
"num_repeats": int,
|
||||||
|
"random_crop": bool,
|
||||||
|
"shuffle_caption": bool,
|
||||||
|
"keep_tokens": int,
|
||||||
|
"keep_tokens_separator": str,
|
||||||
|
"secondary_separator": str,
|
||||||
|
"caption_separator": str,
|
||||||
|
"enable_wildcard": bool,
|
||||||
|
"token_warmup_min": int,
|
||||||
|
"token_warmup_step": Any(float, int),
|
||||||
|
"caption_prefix": str,
|
||||||
|
"caption_suffix": str,
|
||||||
|
}
|
||||||
|
# DO means DropOut
|
||||||
|
DO_SUBSET_ASCENDABLE_SCHEMA = {
|
||||||
|
"caption_dropout_every_n_epochs": int,
|
||||||
|
"caption_dropout_rate": Any(float, int),
|
||||||
|
"caption_tag_dropout_rate": Any(float, int),
|
||||||
|
}
|
||||||
|
# DB means DreamBooth
|
||||||
|
DB_SUBSET_ASCENDABLE_SCHEMA = {
|
||||||
|
"caption_extension": str,
|
||||||
|
"class_tokens": str,
|
||||||
|
"cache_info": bool,
|
||||||
|
}
|
||||||
|
DB_SUBSET_DISTINCT_SCHEMA = {
|
||||||
|
Required("image_dir"): str,
|
||||||
|
"is_reg": bool,
|
||||||
|
"alpha_mask": bool,
|
||||||
|
}
|
||||||
|
# FT means FineTuning
|
||||||
|
FT_SUBSET_DISTINCT_SCHEMA = {
|
||||||
|
Required("metadata_file"): str,
|
||||||
|
"image_dir": str,
|
||||||
|
"alpha_mask": bool,
|
||||||
|
}
|
||||||
|
CN_SUBSET_ASCENDABLE_SCHEMA = {
|
||||||
|
"caption_extension": str,
|
||||||
|
"cache_info": bool,
|
||||||
|
}
|
||||||
|
CN_SUBSET_DISTINCT_SCHEMA = {
|
||||||
|
Required("image_dir"): str,
|
||||||
|
Required("conditioning_data_dir"): str,
|
||||||
|
}
|
||||||
|
|
||||||
|
# datasets schema
|
||||||
|
DATASET_ASCENDABLE_SCHEMA = {
|
||||||
|
"batch_size": int,
|
||||||
|
"bucket_no_upscale": bool,
|
||||||
|
"bucket_reso_steps": int,
|
||||||
|
"enable_bucket": bool,
|
||||||
|
"max_bucket_reso": int,
|
||||||
|
"min_bucket_reso": int,
|
||||||
|
"resolution": functools.partial(__validate_and_convert_scalar_or_twodim.__func__, int),
|
||||||
|
"network_multiplier": float,
|
||||||
|
}
|
||||||
|
|
||||||
|
# options handled by argparse but not handled by user config
|
||||||
|
ARGPARSE_SPECIFIC_SCHEMA = {
|
||||||
|
"debug_dataset": bool,
|
||||||
|
"max_token_length": Any(None, int),
|
||||||
|
"prior_loss_weight": Any(float, int),
|
||||||
|
}
|
||||||
|
# for handling default None value of argparse
|
||||||
|
ARGPARSE_NULLABLE_OPTNAMES = [
|
||||||
|
"face_crop_aug_range",
|
||||||
|
"resolution",
|
||||||
|
]
|
||||||
|
# prepare map because option name may differ among argparse and user config
|
||||||
|
ARGPARSE_OPTNAME_TO_CONFIG_OPTNAME = {
|
||||||
|
"train_batch_size": "batch_size",
|
||||||
|
"dataset_repeats": "num_repeats",
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(self, support_dreambooth: bool, support_finetuning: bool, support_controlnet: bool, support_dropout: bool) -> None:
|
||||||
|
assert support_dreambooth or support_finetuning or support_controlnet, (
|
||||||
|
"Neither DreamBooth mode nor fine tuning mode nor controlnet mode specified. Please specify one mode or more."
|
||||||
|
+ " / DreamBooth モードか fine tuning モードか controlnet モードのどれも指定されていません。1つ以上指定してください。"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.db_subset_schema = self.__merge_dict(
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DB_SUBSET_DISTINCT_SCHEMA,
|
||||||
|
self.DB_SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.ft_subset_schema = self.__merge_dict(
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.FT_SUBSET_DISTINCT_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.cn_subset_schema = self.__merge_dict(
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.CN_SUBSET_DISTINCT_SCHEMA,
|
||||||
|
self.CN_SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.db_dataset_schema = self.__merge_dict(
|
||||||
|
self.DATASET_ASCENDABLE_SCHEMA,
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DB_SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
{"subsets": [self.db_subset_schema]},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.ft_dataset_schema = self.__merge_dict(
|
||||||
|
self.DATASET_ASCENDABLE_SCHEMA,
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
{"subsets": [self.ft_subset_schema]},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.cn_dataset_schema = self.__merge_dict(
|
||||||
|
self.DATASET_ASCENDABLE_SCHEMA,
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.CN_SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
{"subsets": [self.cn_subset_schema]},
|
||||||
|
)
|
||||||
|
|
||||||
|
if support_dreambooth and support_finetuning:
|
||||||
|
|
||||||
|
def validate_flex_dataset(dataset_config: dict):
|
||||||
|
subsets_config = dataset_config.get("subsets", [])
|
||||||
|
|
||||||
|
if support_controlnet and all(["conditioning_data_dir" in subset for subset in subsets_config]):
|
||||||
|
return Schema(self.cn_dataset_schema)(dataset_config)
|
||||||
|
# check dataset meets FT style
|
||||||
|
# NOTE: all FT subsets should have "metadata_file"
|
||||||
|
elif all(["metadata_file" in subset for subset in subsets_config]):
|
||||||
|
return Schema(self.ft_dataset_schema)(dataset_config)
|
||||||
|
# check dataset meets DB style
|
||||||
|
# NOTE: all DB subsets should have no "metadata_file"
|
||||||
|
elif all(["metadata_file" not in subset for subset in subsets_config]):
|
||||||
|
return Schema(self.db_dataset_schema)(dataset_config)
|
||||||
|
else:
|
||||||
|
raise voluptuous.Invalid(
|
||||||
|
"DreamBooth subset and fine tuning subset cannot be mixed in the same dataset. Please split them into separate datasets. / DreamBoothのサブセットとfine tuninのサブセットを同一のデータセットに混在させることはできません。別々のデータセットに分割してください。"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.dataset_schema = validate_flex_dataset
|
||||||
|
elif support_dreambooth:
|
||||||
|
if support_controlnet:
|
||||||
|
self.dataset_schema = self.cn_dataset_schema
|
||||||
|
else:
|
||||||
|
self.dataset_schema = self.db_dataset_schema
|
||||||
|
elif support_finetuning:
|
||||||
|
self.dataset_schema = self.ft_dataset_schema
|
||||||
|
elif support_controlnet:
|
||||||
|
self.dataset_schema = self.cn_dataset_schema
|
||||||
|
|
||||||
|
self.general_schema = self.__merge_dict(
|
||||||
|
self.DATASET_ASCENDABLE_SCHEMA,
|
||||||
|
self.SUBSET_ASCENDABLE_SCHEMA,
|
||||||
|
self.DB_SUBSET_ASCENDABLE_SCHEMA if support_dreambooth else {},
|
||||||
|
self.CN_SUBSET_ASCENDABLE_SCHEMA if support_controlnet else {},
|
||||||
|
self.DO_SUBSET_ASCENDABLE_SCHEMA if support_dropout else {},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.user_config_validator = Schema(
|
||||||
|
{
|
||||||
|
"general": self.general_schema,
|
||||||
|
"datasets": [self.dataset_schema],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.argparse_schema = self.__merge_dict(
|
||||||
|
self.general_schema,
|
||||||
|
self.ARGPARSE_SPECIFIC_SCHEMA,
|
||||||
|
{optname: Any(None, self.general_schema[optname]) for optname in self.ARGPARSE_NULLABLE_OPTNAMES},
|
||||||
|
{a_name: self.general_schema[c_name] for a_name, c_name in self.ARGPARSE_OPTNAME_TO_CONFIG_OPTNAME.items()},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.argparse_config_validator = Schema(Object(self.argparse_schema), extra=voluptuous.ALLOW_EXTRA)
|
||||||
|
|
||||||
|
def sanitize_user_config(self, user_config: dict) -> dict:
|
||||||
|
try:
|
||||||
|
return self.user_config_validator(user_config)
|
||||||
|
except MultipleInvalid:
|
||||||
|
# TODO: エラー発生時のメッセージをわかりやすくする
|
||||||
|
logger.error("Invalid user config / ユーザ設定の形式が正しくないようです")
|
||||||
|
raise
|
||||||
|
|
||||||
|
# NOTE: In nature, argument parser result is not needed to be sanitize
|
||||||
|
# However this will help us to detect program bug
|
||||||
|
def sanitize_argparse_namespace(self, argparse_namespace: argparse.Namespace) -> argparse.Namespace:
|
||||||
|
try:
|
||||||
|
return self.argparse_config_validator(argparse_namespace)
|
||||||
|
except MultipleInvalid:
|
||||||
|
# XXX: this should be a bug
|
||||||
|
logger.error(
|
||||||
|
"Invalid cmdline parsed arguments. This should be a bug. / コマンドラインのパース結果が正しくないようです。プログラムのバグの可能性が高いです。"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
# NOTE: value would be overwritten by latter dict if there is already the same key
|
||||||
|
@staticmethod
|
||||||
|
def __merge_dict(*dict_list: dict) -> dict:
|
||||||
|
merged = {}
|
||||||
|
for schema in dict_list:
|
||||||
|
# merged |= schema
|
||||||
|
for k, v in schema.items():
|
||||||
|
merged[k] = v
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
|
class BlueprintGenerator:
|
||||||
|
BLUEPRINT_PARAM_NAME_TO_CONFIG_OPTNAME = {}
|
||||||
|
|
||||||
|
def __init__(self, sanitizer: ConfigSanitizer):
|
||||||
|
self.sanitizer = sanitizer
|
||||||
|
|
||||||
|
# runtime_params is for parameters which is only configurable on runtime, such as tokenizer
|
||||||
|
def generate(self, user_config: dict, argparse_namespace: argparse.Namespace, **runtime_params) -> Blueprint:
|
||||||
|
sanitized_user_config = self.sanitizer.sanitize_user_config(user_config)
|
||||||
|
sanitized_argparse_namespace = self.sanitizer.sanitize_argparse_namespace(argparse_namespace)
|
||||||
|
|
||||||
|
# convert argparse namespace to dict like config
|
||||||
|
# NOTE: it is ok to have extra entries in dict
|
||||||
|
optname_map = self.sanitizer.ARGPARSE_OPTNAME_TO_CONFIG_OPTNAME
|
||||||
|
argparse_config = {
|
||||||
|
optname_map.get(optname, optname): value for optname, value in vars(sanitized_argparse_namespace).items()
|
||||||
|
}
|
||||||
|
|
||||||
|
general_config = sanitized_user_config.get("general", {})
|
||||||
|
|
||||||
|
dataset_blueprints = []
|
||||||
|
for dataset_config in sanitized_user_config.get("datasets", []):
|
||||||
|
# NOTE: if subsets have no "metadata_file", these are DreamBooth datasets/subsets
|
||||||
|
subsets = dataset_config.get("subsets", [])
|
||||||
|
is_dreambooth = all(["metadata_file" not in subset for subset in subsets])
|
||||||
|
is_controlnet = all(["conditioning_data_dir" in subset for subset in subsets])
|
||||||
|
if is_controlnet:
|
||||||
|
subset_params_klass = ControlNetSubsetParams
|
||||||
|
dataset_params_klass = ControlNetDatasetParams
|
||||||
|
elif is_dreambooth:
|
||||||
|
subset_params_klass = DreamBoothSubsetParams
|
||||||
|
dataset_params_klass = DreamBoothDatasetParams
|
||||||
|
else:
|
||||||
|
subset_params_klass = FineTuningSubsetParams
|
||||||
|
dataset_params_klass = FineTuningDatasetParams
|
||||||
|
|
||||||
|
subset_blueprints = []
|
||||||
|
for subset_config in subsets:
|
||||||
|
params = self.generate_params_by_fallbacks(
|
||||||
|
subset_params_klass, [subset_config, dataset_config, general_config, argparse_config, runtime_params]
|
||||||
|
)
|
||||||
|
subset_blueprints.append(SubsetBlueprint(params))
|
||||||
|
|
||||||
|
params = self.generate_params_by_fallbacks(
|
||||||
|
dataset_params_klass, [dataset_config, general_config, argparse_config, runtime_params]
|
||||||
|
)
|
||||||
|
dataset_blueprints.append(DatasetBlueprint(is_dreambooth, is_controlnet, params, subset_blueprints))
|
||||||
|
|
||||||
|
dataset_group_blueprint = DatasetGroupBlueprint(dataset_blueprints)
|
||||||
|
|
||||||
|
return Blueprint(dataset_group_blueprint)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def generate_params_by_fallbacks(param_klass, fallbacks: Sequence[dict]):
|
||||||
|
name_map = BlueprintGenerator.BLUEPRINT_PARAM_NAME_TO_CONFIG_OPTNAME
|
||||||
|
search_value = BlueprintGenerator.search_value
|
||||||
|
default_params = asdict(param_klass())
|
||||||
|
param_names = default_params.keys()
|
||||||
|
|
||||||
|
params = {name: search_value(name_map.get(name, name), fallbacks, default_params.get(name)) for name in param_names}
|
||||||
|
|
||||||
|
return param_klass(**params)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def search_value(key: str, fallbacks: Sequence[dict], default_value=None):
|
||||||
|
for cand in fallbacks:
|
||||||
|
value = cand.get(key)
|
||||||
|
if value is not None:
|
||||||
|
return value
|
||||||
|
|
||||||
|
return default_value
|
||||||
|
|
||||||
|
|
||||||
|
def generate_dataset_group_by_blueprint(dataset_group_blueprint: DatasetGroupBlueprint):
|
||||||
|
datasets: List[Union[DreamBoothDataset, FineTuningDataset, ControlNetDataset]] = []
|
||||||
|
|
||||||
|
for dataset_blueprint in dataset_group_blueprint.datasets:
|
||||||
|
if dataset_blueprint.is_controlnet:
|
||||||
|
subset_klass = ControlNetSubset
|
||||||
|
dataset_klass = ControlNetDataset
|
||||||
|
elif dataset_blueprint.is_dreambooth:
|
||||||
|
subset_klass = DreamBoothSubset
|
||||||
|
dataset_klass = DreamBoothDataset
|
||||||
|
else:
|
||||||
|
subset_klass = FineTuningSubset
|
||||||
|
dataset_klass = FineTuningDataset
|
||||||
|
|
||||||
|
subsets = [subset_klass(**asdict(subset_blueprint.params)) for subset_blueprint in dataset_blueprint.subsets]
|
||||||
|
dataset = dataset_klass(subsets=subsets, **asdict(dataset_blueprint.params))
|
||||||
|
datasets.append(dataset)
|
||||||
|
|
||||||
|
# print info
|
||||||
|
info = ""
|
||||||
|
for i, dataset in enumerate(datasets):
|
||||||
|
is_dreambooth = isinstance(dataset, DreamBoothDataset)
|
||||||
|
is_controlnet = isinstance(dataset, ControlNetDataset)
|
||||||
|
info += dedent(
|
||||||
|
f"""\
|
||||||
|
[Dataset {i}]
|
||||||
|
batch_size: {dataset.batch_size}
|
||||||
|
resolution: {(dataset.width, dataset.height)}
|
||||||
|
enable_bucket: {dataset.enable_bucket}
|
||||||
|
network_multiplier: {dataset.network_multiplier}
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
if dataset.enable_bucket:
|
||||||
|
info += indent(
|
||||||
|
dedent(
|
||||||
|
f"""\
|
||||||
|
min_bucket_reso: {dataset.min_bucket_reso}
|
||||||
|
max_bucket_reso: {dataset.max_bucket_reso}
|
||||||
|
bucket_reso_steps: {dataset.bucket_reso_steps}
|
||||||
|
bucket_no_upscale: {dataset.bucket_no_upscale}
|
||||||
|
\n"""
|
||||||
|
),
|
||||||
|
" ",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
info += "\n"
|
||||||
|
|
||||||
|
for j, subset in enumerate(dataset.subsets):
|
||||||
|
info += indent(
|
||||||
|
dedent(
|
||||||
|
f"""\
|
||||||
|
[Subset {j} of Dataset {i}]
|
||||||
|
image_dir: "{subset.image_dir}"
|
||||||
|
image_count: {subset.img_count}
|
||||||
|
num_repeats: {subset.num_repeats}
|
||||||
|
shuffle_caption: {subset.shuffle_caption}
|
||||||
|
keep_tokens: {subset.keep_tokens}
|
||||||
|
keep_tokens_separator: {subset.keep_tokens_separator}
|
||||||
|
caption_separator: {subset.caption_separator}
|
||||||
|
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_tag_dropout_rate: {subset.caption_tag_dropout_rate}
|
||||||
|
caption_prefix: {subset.caption_prefix}
|
||||||
|
caption_suffix: {subset.caption_suffix}
|
||||||
|
color_aug: {subset.color_aug}
|
||||||
|
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},
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
" ",
|
||||||
|
)
|
||||||
|
|
||||||
|
if is_dreambooth:
|
||||||
|
info += indent(
|
||||||
|
dedent(
|
||||||
|
f"""\
|
||||||
|
is_reg: {subset.is_reg}
|
||||||
|
class_tokens: {subset.class_tokens}
|
||||||
|
caption_extension: {subset.caption_extension}
|
||||||
|
\n"""
|
||||||
|
),
|
||||||
|
" ",
|
||||||
|
)
|
||||||
|
elif not is_controlnet:
|
||||||
|
info += indent(
|
||||||
|
dedent(
|
||||||
|
f"""\
|
||||||
|
metadata_file: {subset.metadata_file}
|
||||||
|
\n"""
|
||||||
|
),
|
||||||
|
" ",
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(f"{info}")
|
||||||
|
|
||||||
|
# make buckets first because it determines the length of dataset
|
||||||
|
# and set the same seed for all datasets
|
||||||
|
seed = random.randint(0, 2**31) # actual seed is seed + epoch_no
|
||||||
|
for i, dataset in enumerate(datasets):
|
||||||
|
logger.info(f"[Dataset {i}]")
|
||||||
|
dataset.make_buckets()
|
||||||
|
dataset.set_seed(seed)
|
||||||
|
|
||||||
|
return DatasetGroup(datasets)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_dreambooth_subsets_config_by_subdirs(train_data_dir: Optional[str] = None, reg_data_dir: Optional[str] = None):
|
||||||
|
def extract_dreambooth_params(name: str) -> Tuple[int, str]:
|
||||||
|
tokens = name.split("_")
|
||||||
|
try:
|
||||||
|
n_repeats = int(tokens[0])
|
||||||
|
except ValueError as e:
|
||||||
|
logger.warning(f"ignore directory without repeats / 繰り返し回数のないディレクトリを無視します: {name}")
|
||||||
|
return 0, ""
|
||||||
|
caption_by_folder = "_".join(tokens[1:])
|
||||||
|
return n_repeats, caption_by_folder
|
||||||
|
|
||||||
|
def generate(base_dir: Optional[str], is_reg: bool):
|
||||||
|
if base_dir is None:
|
||||||
|
return []
|
||||||
|
|
||||||
|
base_dir: Path = Path(base_dir)
|
||||||
|
if not base_dir.is_dir():
|
||||||
|
return []
|
||||||
|
|
||||||
|
subsets_config = []
|
||||||
|
for subdir in base_dir.iterdir():
|
||||||
|
if not subdir.is_dir():
|
||||||
|
continue
|
||||||
|
|
||||||
|
num_repeats, class_tokens = extract_dreambooth_params(subdir.name)
|
||||||
|
if num_repeats < 1:
|
||||||
|
continue
|
||||||
|
|
||||||
|
subset_config = {"image_dir": str(subdir), "num_repeats": num_repeats, "is_reg": is_reg, "class_tokens": class_tokens}
|
||||||
|
subsets_config.append(subset_config)
|
||||||
|
|
||||||
|
return subsets_config
|
||||||
|
|
||||||
|
subsets_config = []
|
||||||
|
subsets_config += generate(train_data_dir, False)
|
||||||
|
subsets_config += generate(reg_data_dir, True)
|
||||||
|
|
||||||
|
return subsets_config
|
||||||
|
|
||||||
|
|
||||||
|
def generate_controlnet_subsets_config_by_subdirs(
|
||||||
|
train_data_dir: Optional[str] = None, conditioning_data_dir: Optional[str] = None, caption_extension: str = ".txt"
|
||||||
|
):
|
||||||
|
def generate(base_dir: Optional[str]):
|
||||||
|
if base_dir is None:
|
||||||
|
return []
|
||||||
|
|
||||||
|
base_dir: Path = Path(base_dir)
|
||||||
|
if not base_dir.is_dir():
|
||||||
|
return []
|
||||||
|
|
||||||
|
subsets_config = []
|
||||||
|
subset_config = {
|
||||||
|
"image_dir": train_data_dir,
|
||||||
|
"conditioning_data_dir": conditioning_data_dir,
|
||||||
|
"caption_extension": caption_extension,
|
||||||
|
"num_repeats": 1,
|
||||||
|
}
|
||||||
|
subsets_config.append(subset_config)
|
||||||
|
|
||||||
|
return subsets_config
|
||||||
|
|
||||||
|
subsets_config = []
|
||||||
|
subsets_config += generate(train_data_dir)
|
||||||
|
|
||||||
|
return subsets_config
|
||||||
|
|
||||||
|
|
||||||
|
def load_user_config(file: str) -> dict:
|
||||||
|
file_path: Path = Path(file)
|
||||||
|
if not file_path.is_file():
|
||||||
|
#raise ValueError(f"file not found / ファイルが見つかりません: {file}")
|
||||||
|
return toml.loads(file)
|
||||||
|
|
||||||
|
if file_path.name.lower().endswith(".json"):
|
||||||
|
try:
|
||||||
|
with open(file, "r") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
except Exception:
|
||||||
|
logger.error(
|
||||||
|
f"Error on parsing JSON config file. Please check the format. / JSON 形式の設定ファイルの読み込みに失敗しました。文法が正しいか確認してください。: {file}"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
elif file_path.name.lower().endswith(".toml"):
|
||||||
|
try:
|
||||||
|
config = toml.load(file_path)
|
||||||
|
except Exception:
|
||||||
|
logger.error(
|
||||||
|
f"Error on parsing TOML config file. Please check the format. / TOML 形式の設定ファイルの読み込みに失敗しました。文法が正しいか確認してください。: {file}"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
else:
|
||||||
|
raise ValueError(f"not supported config file format / 対応していない設定ファイルの形式です: {file_path}")
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
# for config test
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--support_dreambooth", action="store_true")
|
||||||
|
parser.add_argument("--support_finetuning", action="store_true")
|
||||||
|
parser.add_argument("--support_controlnet", action="store_true")
|
||||||
|
parser.add_argument("--support_dropout", action="store_true")
|
||||||
|
parser.add_argument("dataset_config")
|
||||||
|
config_args, remain = parser.parse_known_args()
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
train_util.add_dataset_arguments(
|
||||||
|
parser, config_args.support_dreambooth, config_args.support_finetuning, config_args.support_dropout
|
||||||
|
)
|
||||||
|
train_util.add_training_arguments(parser, config_args.support_dreambooth)
|
||||||
|
argparse_namespace = parser.parse_args(remain)
|
||||||
|
train_util.prepare_dataset_args(argparse_namespace, config_args.support_finetuning)
|
||||||
|
|
||||||
|
logger.info("[argparse_namespace]")
|
||||||
|
logger.info(f"{vars(argparse_namespace)}")
|
||||||
|
|
||||||
|
user_config = load_user_config(config_args.dataset_config)
|
||||||
|
|
||||||
|
logger.info("")
|
||||||
|
logger.info("[user_config]")
|
||||||
|
logger.info(f"{user_config}")
|
||||||
|
|
||||||
|
sanitizer = ConfigSanitizer(
|
||||||
|
config_args.support_dreambooth, config_args.support_finetuning, config_args.support_controlnet, config_args.support_dropout
|
||||||
|
)
|
||||||
|
sanitized_user_config = sanitizer.sanitize_user_config(user_config)
|
||||||
|
|
||||||
|
logger.info("")
|
||||||
|
logger.info("[sanitized_user_config]")
|
||||||
|
logger.info(f"{sanitized_user_config}")
|
||||||
|
|
||||||
|
blueprint = BlueprintGenerator(sanitizer).generate(user_config, argparse_namespace)
|
||||||
|
|
||||||
|
logger.info("")
|
||||||
|
logger.info("[blueprint]")
|
||||||
|
logger.info(f"{blueprint}")
|
||||||
@@ -0,0 +1,556 @@
|
|||||||
|
import torch
|
||||||
|
import argparse
|
||||||
|
import random
|
||||||
|
import re
|
||||||
|
from typing import List, Optional, Union
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_scheduler_for_custom_training(noise_scheduler, device):
|
||||||
|
if hasattr(noise_scheduler, "all_snr"):
|
||||||
|
return
|
||||||
|
|
||||||
|
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||||
|
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||||
|
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||||
|
alpha = sqrt_alphas_cumprod
|
||||||
|
sigma = sqrt_one_minus_alphas_cumprod
|
||||||
|
all_snr = (alpha / sigma) ** 2
|
||||||
|
|
||||||
|
noise_scheduler.all_snr = all_snr.to(device)
|
||||||
|
|
||||||
|
|
||||||
|
def fix_noise_scheduler_betas_for_zero_terminal_snr(noise_scheduler):
|
||||||
|
# fix beta: zero terminal SNR
|
||||||
|
logger.info(f"fix noise scheduler betas: https://arxiv.org/abs/2305.08891")
|
||||||
|
|
||||||
|
def enforce_zero_terminal_snr(betas):
|
||||||
|
# Convert betas to alphas_bar_sqrt
|
||||||
|
alphas = 1 - betas
|
||||||
|
alphas_bar = alphas.cumprod(0)
|
||||||
|
alphas_bar_sqrt = alphas_bar.sqrt()
|
||||||
|
|
||||||
|
# Store old values.
|
||||||
|
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||||
|
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||||
|
# Shift so last timestep is zero.
|
||||||
|
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||||
|
# Scale so first timestep is back to old value.
|
||||||
|
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||||
|
|
||||||
|
# Convert alphas_bar_sqrt to betas
|
||||||
|
alphas_bar = alphas_bar_sqrt**2
|
||||||
|
alphas = alphas_bar[1:] / alphas_bar[:-1]
|
||||||
|
alphas = torch.cat([alphas_bar[0:1], alphas])
|
||||||
|
betas = 1 - alphas
|
||||||
|
return betas
|
||||||
|
|
||||||
|
betas = noise_scheduler.betas
|
||||||
|
betas = enforce_zero_terminal_snr(betas)
|
||||||
|
alphas = 1.0 - betas
|
||||||
|
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||||
|
|
||||||
|
# logger.info(f"original: {noise_scheduler.betas}")
|
||||||
|
# logger.info(f"fixed: {betas}")
|
||||||
|
|
||||||
|
noise_scheduler.betas = betas
|
||||||
|
noise_scheduler.alphas = alphas
|
||||||
|
noise_scheduler.alphas_cumprod = alphas_cumprod
|
||||||
|
|
||||||
|
|
||||||
|
def apply_snr_weight(loss, timesteps, noise_scheduler, gamma, v_prediction=False):
|
||||||
|
snr = torch.stack([noise_scheduler.all_snr[t] for t in timesteps])
|
||||||
|
min_snr_gamma = torch.minimum(snr, torch.full_like(snr, gamma))
|
||||||
|
if v_prediction:
|
||||||
|
snr_weight = torch.div(min_snr_gamma, snr + 1).float().to(loss.device)
|
||||||
|
else:
|
||||||
|
snr_weight = torch.div(min_snr_gamma, snr).float().to(loss.device)
|
||||||
|
loss = loss * snr_weight
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
def scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler):
|
||||||
|
scale = get_snr_scale(timesteps, noise_scheduler)
|
||||||
|
loss = loss * scale
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
def get_snr_scale(timesteps, noise_scheduler):
|
||||||
|
snr_t = torch.stack([noise_scheduler.all_snr[t] for t in timesteps]) # batch_size
|
||||||
|
snr_t = torch.minimum(snr_t, torch.ones_like(snr_t) * 1000) # if timestep is 0, snr_t is inf, so limit it to 1000
|
||||||
|
scale = snr_t / (snr_t + 1)
|
||||||
|
# # show debug info
|
||||||
|
# logger.info(f"timesteps: {timesteps}, snr_t: {snr_t}, scale: {scale}")
|
||||||
|
return scale
|
||||||
|
|
||||||
|
|
||||||
|
def add_v_prediction_like_loss(loss, timesteps, noise_scheduler, v_pred_like_loss):
|
||||||
|
scale = get_snr_scale(timesteps, noise_scheduler)
|
||||||
|
# logger.info(f"add v-prediction like loss: {v_pred_like_loss}, scale: {scale}, loss: {loss}, time: {timesteps}")
|
||||||
|
loss = loss + loss / scale * v_pred_like_loss
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
def apply_debiased_estimation(loss, timesteps, noise_scheduler):
|
||||||
|
snr_t = torch.stack([noise_scheduler.all_snr[t] for t in timesteps]) # batch_size
|
||||||
|
snr_t = torch.minimum(snr_t, torch.ones_like(snr_t) * 1000) # if timestep is 0, snr_t is inf, so limit it to 1000
|
||||||
|
weight = 1 / torch.sqrt(snr_t)
|
||||||
|
loss = weight * loss
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
# TODO train_utilと分散しているのでどちらかに寄せる
|
||||||
|
|
||||||
|
|
||||||
|
def add_custom_train_arguments(parser: argparse.ArgumentParser, support_weighted_captions: bool = True):
|
||||||
|
parser.add_argument(
|
||||||
|
"--min_snr_gamma",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="gamma for reducing the weight of high loss timesteps. Lower numbers have stronger effect. 5 is recommended by paper. / 低いタイムステップでの高いlossに対して重みを減らすためのgamma値、低いほど効果が強く、論文では5が推奨",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--scale_v_pred_loss_like_noise_pred",
|
||||||
|
action="store_true",
|
||||||
|
help="scale v-prediction loss like noise prediction loss / v-prediction lossをnoise prediction lossと同じようにスケーリングする",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--v_pred_like_loss",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="add v-prediction like loss multiplied by this value / v-prediction lossをこの値をかけたものをlossに加算する",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--debiased_estimation_loss",
|
||||||
|
action="store_true",
|
||||||
|
help="debiased estimation loss / debiased estimation loss",
|
||||||
|
)
|
||||||
|
if support_weighted_captions:
|
||||||
|
parser.add_argument(
|
||||||
|
"--weighted_captions",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help="Enable weighted captions in the standard style (token:1.3). No commas inside parens, or shuffle/dropout may break the decoder. / 「[token]」、「(token)」「(token:1.3)」のような重み付きキャプションを有効にする。カンマを括弧内に入れるとシャッフルやdropoutで重みづけがおかしくなるので注意",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
re_attention = re.compile(
|
||||||
|
r"""
|
||||||
|
\\\(|
|
||||||
|
\\\)|
|
||||||
|
\\\[|
|
||||||
|
\\]|
|
||||||
|
\\\\|
|
||||||
|
\\|
|
||||||
|
\(|
|
||||||
|
\[|
|
||||||
|
:([+-]?[.\d]+)\)|
|
||||||
|
\)|
|
||||||
|
]|
|
||||||
|
[^\\()\[\]:]+|
|
||||||
|
:
|
||||||
|
""",
|
||||||
|
re.X,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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 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(tokenizer, prompt: List[str], max_length: int):
|
||||||
|
r"""
|
||||||
|
Tokenize a list of prompts and return its tokens with weights of each token.
|
||||||
|
|
||||||
|
No padding, starting or ending token is included.
|
||||||
|
"""
|
||||||
|
tokens = []
|
||||||
|
weights = []
|
||||||
|
truncated = False
|
||||||
|
for text in prompt:
|
||||||
|
texts_and_weights = parse_prompt_attention(text)
|
||||||
|
text_token = []
|
||||||
|
text_weight = []
|
||||||
|
for word, weight in texts_and_weights:
|
||||||
|
# tokenize and discard the starting and the ending token
|
||||||
|
token = tokenizer(word).input_ids[1:-1]
|
||||||
|
text_token += token
|
||||||
|
# copy the weight by length of token
|
||||||
|
text_weight += [weight] * len(token)
|
||||||
|
# stop if the text is too long (longer than truncation limit)
|
||||||
|
if len(text_token) > max_length:
|
||||||
|
truncated = True
|
||||||
|
break
|
||||||
|
# truncate
|
||||||
|
if len(text_token) > max_length:
|
||||||
|
truncated = True
|
||||||
|
text_token = text_token[:max_length]
|
||||||
|
text_weight = text_weight[:max_length]
|
||||||
|
tokens.append(text_token)
|
||||||
|
weights.append(text_weight)
|
||||||
|
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, no_boseos_middle=True, chunk_length=77):
|
||||||
|
r"""
|
||||||
|
Pad the tokens (with starting and ending tokens) and weights (with 1.0) to max_length.
|
||||||
|
"""
|
||||||
|
max_embeddings_multiples = (max_length - 2) // (chunk_length - 2)
|
||||||
|
weights_length = max_length if no_boseos_middle else max_embeddings_multiples * chunk_length
|
||||||
|
for i in range(len(tokens)):
|
||||||
|
tokens[i] = [bos] + tokens[i] + [eos] * (max_length - 1 - len(tokens[i]))
|
||||||
|
if no_boseos_middle:
|
||||||
|
weights[i] = [1.0] + weights[i] + [1.0] * (max_length - 1 - len(weights[i]))
|
||||||
|
else:
|
||||||
|
w = []
|
||||||
|
if len(weights[i]) == 0:
|
||||||
|
w = [1.0] * weights_length
|
||||||
|
else:
|
||||||
|
for j in range(max_embeddings_multiples):
|
||||||
|
w.append(1.0) # weight for starting token in this chunk
|
||||||
|
w += weights[i][j * (chunk_length - 2) : min(len(weights[i]), (j + 1) * (chunk_length - 2))]
|
||||||
|
w.append(1.0) # weight for ending token in this chunk
|
||||||
|
w += [1.0] * (weights_length - len(w))
|
||||||
|
weights[i] = w[:]
|
||||||
|
|
||||||
|
return tokens, weights
|
||||||
|
|
||||||
|
|
||||||
|
def get_unweighted_text_embeddings(
|
||||||
|
tokenizer,
|
||||||
|
text_encoder,
|
||||||
|
text_input: torch.Tensor,
|
||||||
|
chunk_length: int,
|
||||||
|
clip_skip: int,
|
||||||
|
eos: int,
|
||||||
|
pad: int,
|
||||||
|
no_boseos_middle: Optional[bool] = True,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
When the length of tokens is a multiple of the capacity of the text encoder,
|
||||||
|
it should be split into chunks and sent to the text encoder individually.
|
||||||
|
"""
|
||||||
|
max_embeddings_multiples = (text_input.shape[1] - 2) // (chunk_length - 2)
|
||||||
|
if max_embeddings_multiples > 1:
|
||||||
|
text_embeddings = []
|
||||||
|
for i in range(max_embeddings_multiples):
|
||||||
|
# extract the i-th chunk
|
||||||
|
text_input_chunk = text_input[:, i * (chunk_length - 2) : (i + 1) * (chunk_length - 2) + 2].clone()
|
||||||
|
|
||||||
|
# cover the head and the tail by the starting and the ending tokens
|
||||||
|
text_input_chunk[:, 0] = text_input[0, 0]
|
||||||
|
if pad == eos: # v1
|
||||||
|
text_input_chunk[:, -1] = text_input[0, -1]
|
||||||
|
else: # v2
|
||||||
|
for j in range(len(text_input_chunk)):
|
||||||
|
if text_input_chunk[j, -1] != eos and text_input_chunk[j, -1] != pad: # 最後に普通の文字がある
|
||||||
|
text_input_chunk[j, -1] = eos
|
||||||
|
if text_input_chunk[j, 1] == pad: # BOSだけであとはPAD
|
||||||
|
text_input_chunk[j, 1] = eos
|
||||||
|
|
||||||
|
if clip_skip is None or clip_skip == 1:
|
||||||
|
text_embedding = text_encoder(text_input_chunk)[0]
|
||||||
|
else:
|
||||||
|
enc_out = text_encoder(text_input_chunk, output_hidden_states=True, return_dict=True)
|
||||||
|
text_embedding = enc_out["hidden_states"][-clip_skip]
|
||||||
|
text_embedding = text_encoder.text_model.final_layer_norm(text_embedding)
|
||||||
|
|
||||||
|
if no_boseos_middle:
|
||||||
|
if i == 0:
|
||||||
|
# discard the ending token
|
||||||
|
text_embedding = text_embedding[:, :-1]
|
||||||
|
elif i == max_embeddings_multiples - 1:
|
||||||
|
# discard the starting token
|
||||||
|
text_embedding = text_embedding[:, 1:]
|
||||||
|
else:
|
||||||
|
# discard both starting and ending tokens
|
||||||
|
text_embedding = text_embedding[:, 1:-1]
|
||||||
|
|
||||||
|
text_embeddings.append(text_embedding)
|
||||||
|
text_embeddings = torch.concat(text_embeddings, axis=1)
|
||||||
|
else:
|
||||||
|
if clip_skip is None or clip_skip == 1:
|
||||||
|
text_embeddings = text_encoder(text_input)[0]
|
||||||
|
else:
|
||||||
|
enc_out = text_encoder(text_input, output_hidden_states=True, return_dict=True)
|
||||||
|
text_embeddings = enc_out["hidden_states"][-clip_skip]
|
||||||
|
text_embeddings = text_encoder.text_model.final_layer_norm(text_embeddings)
|
||||||
|
return text_embeddings
|
||||||
|
|
||||||
|
|
||||||
|
def get_weighted_text_embeddings(
|
||||||
|
tokenizer,
|
||||||
|
text_encoder,
|
||||||
|
prompt: Union[str, List[str]],
|
||||||
|
device,
|
||||||
|
max_embeddings_multiples: Optional[int] = 3,
|
||||||
|
no_boseos_middle: Optional[bool] = False,
|
||||||
|
clip_skip=None,
|
||||||
|
):
|
||||||
|
r"""
|
||||||
|
Prompts can be assigned with local weights using brackets. For example,
|
||||||
|
prompt 'A (very beautiful) masterpiece' highlights the words 'very beautiful',
|
||||||
|
and the embedding tokens corresponding to the words get multiplied by a constant, 1.1.
|
||||||
|
|
||||||
|
Also, to regularize of the embedding, the weighted embedding would be scaled to preserve the original mean.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt (`str` or `List[str]`):
|
||||||
|
The prompt or prompts to guide the image generation.
|
||||||
|
max_embeddings_multiples (`int`, *optional*, defaults to `3`):
|
||||||
|
The max multiple length of prompt embeddings compared to the max output length of text encoder.
|
||||||
|
no_boseos_middle (`bool`, *optional*, defaults to `False`):
|
||||||
|
If the length of text token is multiples of the capacity of text encoder, whether reserve the starting and
|
||||||
|
ending token in each of the chunk in the middle.
|
||||||
|
skip_parsing (`bool`, *optional*, defaults to `False`):
|
||||||
|
Skip the parsing of brackets.
|
||||||
|
skip_weighting (`bool`, *optional*, defaults to `False`):
|
||||||
|
Skip the weighting. When the parsing is skipped, it is forced True.
|
||||||
|
"""
|
||||||
|
max_length = (tokenizer.model_max_length - 2) * max_embeddings_multiples + 2
|
||||||
|
if isinstance(prompt, str):
|
||||||
|
prompt = [prompt]
|
||||||
|
|
||||||
|
prompt_tokens, prompt_weights = get_prompts_with_weights(tokenizer, prompt, max_length - 2)
|
||||||
|
|
||||||
|
# round up the longest length of tokens to a multiple of (model_max_length - 2)
|
||||||
|
max_length = max([len(token) for token in prompt_tokens])
|
||||||
|
|
||||||
|
max_embeddings_multiples = min(
|
||||||
|
max_embeddings_multiples,
|
||||||
|
(max_length - 1) // (tokenizer.model_max_length - 2) + 1,
|
||||||
|
)
|
||||||
|
max_embeddings_multiples = max(1, max_embeddings_multiples)
|
||||||
|
max_length = (tokenizer.model_max_length - 2) * max_embeddings_multiples + 2
|
||||||
|
|
||||||
|
# pad the length of tokens and weights
|
||||||
|
bos = tokenizer.bos_token_id
|
||||||
|
eos = tokenizer.eos_token_id
|
||||||
|
pad = tokenizer.pad_token_id
|
||||||
|
prompt_tokens, prompt_weights = pad_tokens_and_weights(
|
||||||
|
prompt_tokens,
|
||||||
|
prompt_weights,
|
||||||
|
max_length,
|
||||||
|
bos,
|
||||||
|
eos,
|
||||||
|
no_boseos_middle=no_boseos_middle,
|
||||||
|
chunk_length=tokenizer.model_max_length,
|
||||||
|
)
|
||||||
|
prompt_tokens = torch.tensor(prompt_tokens, dtype=torch.long, device=device)
|
||||||
|
|
||||||
|
# get the embeddings
|
||||||
|
text_embeddings = get_unweighted_text_embeddings(
|
||||||
|
tokenizer,
|
||||||
|
text_encoder,
|
||||||
|
prompt_tokens,
|
||||||
|
tokenizer.model_max_length,
|
||||||
|
clip_skip,
|
||||||
|
eos,
|
||||||
|
pad,
|
||||||
|
no_boseos_middle=no_boseos_middle,
|
||||||
|
)
|
||||||
|
prompt_weights = torch.tensor(prompt_weights, dtype=text_embeddings.dtype, device=device)
|
||||||
|
|
||||||
|
# assign weights to the prompts and normalize in the sense of mean
|
||||||
|
previous_mean = text_embeddings.float().mean(axis=[-2, -1]).to(text_embeddings.dtype)
|
||||||
|
text_embeddings = text_embeddings * prompt_weights.unsqueeze(-1)
|
||||||
|
current_mean = text_embeddings.float().mean(axis=[-2, -1]).to(text_embeddings.dtype)
|
||||||
|
text_embeddings = text_embeddings * (previous_mean / current_mean).unsqueeze(-1).unsqueeze(-1)
|
||||||
|
|
||||||
|
return text_embeddings
|
||||||
|
|
||||||
|
|
||||||
|
# https://wandb.ai/johnowhitaker/multires_noise/reports/Multi-Resolution-Noise-for-Diffusion-Model-Training--VmlldzozNjYyOTU2
|
||||||
|
def pyramid_noise_like(noise, device, iterations=6, discount=0.4):
|
||||||
|
b, c, w, h = noise.shape # EDIT: w and h get over-written, rename for a different variant!
|
||||||
|
u = torch.nn.Upsample(size=(w, h), mode="bilinear").to(device)
|
||||||
|
for i in range(iterations):
|
||||||
|
r = random.random() * 2 + 2 # Rather than always going 2x,
|
||||||
|
wn, hn = max(1, int(w / (r**i))), max(1, int(h / (r**i)))
|
||||||
|
noise += u(torch.randn(b, c, wn, hn).to(device)) * discount**i
|
||||||
|
if wn == 1 or hn == 1:
|
||||||
|
break # Lowest resolution is 1x1
|
||||||
|
return noise / noise.std() # Scaled back to roughly unit variance
|
||||||
|
|
||||||
|
|
||||||
|
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
|
||||||
|
def apply_noise_offset(latents, noise, noise_offset, adaptive_noise_scale):
|
||||||
|
if noise_offset is None:
|
||||||
|
return noise
|
||||||
|
if adaptive_noise_scale is not None:
|
||||||
|
# latent shape: (batch_size, channels, height, width)
|
||||||
|
# abs mean value for each channel
|
||||||
|
latent_mean = torch.abs(latents.mean(dim=(2, 3), keepdim=True))
|
||||||
|
|
||||||
|
# multiply adaptive noise scale to the mean value and add it to the noise offset
|
||||||
|
noise_offset = noise_offset + adaptive_noise_scale * latent_mean
|
||||||
|
noise_offset = torch.clamp(noise_offset, 0.0, None) # in case of adaptive noise scale is negative
|
||||||
|
|
||||||
|
noise = noise + noise_offset * torch.randn((latents.shape[0], latents.shape[1], 1, 1), device=latents.device)
|
||||||
|
return noise
|
||||||
|
|
||||||
|
|
||||||
|
def apply_masked_loss(loss, batch):
|
||||||
|
if "conditioning_images" in batch:
|
||||||
|
# conditioning image is -1 to 1. we need to convert it to 0 to 1
|
||||||
|
mask_image = batch["conditioning_images"].to(dtype=loss.dtype)[:, 0].unsqueeze(1) # use R channel
|
||||||
|
mask_image = mask_image / 2 + 0.5
|
||||||
|
# print(f"conditioning_image: {mask_image.shape}")
|
||||||
|
elif "alpha_masks" in batch and batch["alpha_masks"] is not None:
|
||||||
|
# alpha mask is 0 to 1
|
||||||
|
mask_image = batch["alpha_masks"].to(dtype=loss.dtype).unsqueeze(1) # add channel dimension
|
||||||
|
# print(f"mask_image: {mask_image.shape}, {mask_image.mean()}")
|
||||||
|
else:
|
||||||
|
return loss
|
||||||
|
|
||||||
|
# resize to the same size as the loss
|
||||||
|
mask_image = torch.nn.functional.interpolate(mask_image, size=loss.shape[2:], mode="area")
|
||||||
|
loss = loss * mask_image
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
##########################################
|
||||||
|
# Perlin Noise
|
||||||
|
def rand_perlin_2d(device, shape, res, fade=lambda t: 6 * t**5 - 15 * t**4 + 10 * t**3):
|
||||||
|
delta = (res[0] / shape[0], res[1] / shape[1])
|
||||||
|
d = (shape[0] // res[0], shape[1] // res[1])
|
||||||
|
|
||||||
|
grid = (
|
||||||
|
torch.stack(
|
||||||
|
torch.meshgrid(torch.arange(0, res[0], delta[0], device=device), torch.arange(0, res[1], delta[1], device=device)),
|
||||||
|
dim=-1,
|
||||||
|
)
|
||||||
|
% 1
|
||||||
|
)
|
||||||
|
angles = 2 * torch.pi * torch.rand(res[0] + 1, res[1] + 1, device=device)
|
||||||
|
gradients = torch.stack((torch.cos(angles), torch.sin(angles)), dim=-1)
|
||||||
|
|
||||||
|
tile_grads = (
|
||||||
|
lambda slice1, slice2: gradients[slice1[0] : slice1[1], slice2[0] : slice2[1]]
|
||||||
|
.repeat_interleave(d[0], 0)
|
||||||
|
.repeat_interleave(d[1], 1)
|
||||||
|
)
|
||||||
|
dot = lambda grad, shift: (
|
||||||
|
torch.stack((grid[: shape[0], : shape[1], 0] + shift[0], grid[: shape[0], : shape[1], 1] + shift[1]), dim=-1)
|
||||||
|
* grad[: shape[0], : shape[1]]
|
||||||
|
).sum(dim=-1)
|
||||||
|
|
||||||
|
n00 = dot(tile_grads([0, -1], [0, -1]), [0, 0])
|
||||||
|
n10 = dot(tile_grads([1, None], [0, -1]), [-1, 0])
|
||||||
|
n01 = dot(tile_grads([0, -1], [1, None]), [0, -1])
|
||||||
|
n11 = dot(tile_grads([1, None], [1, None]), [-1, -1])
|
||||||
|
t = fade(grid[: shape[0], : shape[1]])
|
||||||
|
return 1.414 * torch.lerp(torch.lerp(n00, n10, t[..., 0]), torch.lerp(n01, n11, t[..., 0]), t[..., 1])
|
||||||
|
|
||||||
|
|
||||||
|
def rand_perlin_2d_octaves(device, shape, res, octaves=1, persistence=0.5):
|
||||||
|
noise = torch.zeros(shape, device=device)
|
||||||
|
frequency = 1
|
||||||
|
amplitude = 1
|
||||||
|
for _ in range(octaves):
|
||||||
|
noise += amplitude * rand_perlin_2d(device, shape, (frequency * res[0], frequency * res[1]))
|
||||||
|
frequency *= 2
|
||||||
|
amplitude *= persistence
|
||||||
|
return noise
|
||||||
|
|
||||||
|
|
||||||
|
def perlin_noise(noise, device, octaves):
|
||||||
|
_, c, w, h = noise.shape
|
||||||
|
perlin = lambda: rand_perlin_2d_octaves(device, (w, h), (4, 4), octaves)
|
||||||
|
noise_perlin = []
|
||||||
|
for _ in range(c):
|
||||||
|
noise_perlin.append(perlin())
|
||||||
|
noise_perlin = torch.stack(noise_perlin).unsqueeze(0) # (1, c, w, h)
|
||||||
|
noise += noise_perlin # broadcast for each batch
|
||||||
|
return noise / noise.std() # Scaled back to roughly unit variance
|
||||||
|
"""
|
||||||
@@ -0,0 +1,139 @@
|
|||||||
|
import os
|
||||||
|
import argparse
|
||||||
|
import torch
|
||||||
|
from accelerate import DeepSpeedPlugin, Accelerator
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def add_deepspeed_arguments(parser: argparse.ArgumentParser):
|
||||||
|
# DeepSpeed Arguments. https://huggingface.co/docs/accelerate/usage_guides/deepspeed
|
||||||
|
parser.add_argument("--deepspeed", action="store_true", help="enable deepspeed training")
|
||||||
|
parser.add_argument("--zero_stage", type=int, default=2, choices=[0, 1, 2, 3], help="Possible options are 0,1,2,3.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--offload_optimizer_device",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=[None, "cpu", "nvme"],
|
||||||
|
help="Possible options are none|cpu|nvme. Only applicable with ZeRO Stages 2 and 3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--offload_optimizer_nvme_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--offload_param_device",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=[None, "cpu", "nvme"],
|
||||||
|
help="Possible options are none|cpu|nvme. Only applicable with ZeRO Stage 3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--offload_param_nvme_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Possible options are /nvme|/local_nvme. Only applicable with ZeRO Stage 3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--zero3_init_flag",
|
||||||
|
action="store_true",
|
||||||
|
help="Flag to indicate whether to enable `deepspeed.zero.Init` for constructing massive models."
|
||||||
|
"Only applicable with ZeRO Stage-3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--zero3_save_16bit_model",
|
||||||
|
action="store_true",
|
||||||
|
help="Flag to indicate whether to save 16-bit model. Only applicable with ZeRO Stage-3.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--fp16_master_weights_and_gradients",
|
||||||
|
action="store_true",
|
||||||
|
help="fp16_master_and_gradients requires optimizer to support keeping fp16 master and gradients while keeping the optimizer states in fp32.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_deepspeed_args(args: argparse.Namespace):
|
||||||
|
if not args.deepspeed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# To avoid RuntimeError: DataLoader worker exited unexpectedly with exit code 1.
|
||||||
|
args.max_data_loader_n_workers = 1
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_deepspeed_plugin(args: argparse.Namespace):
|
||||||
|
if not args.deepspeed:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
import deepspeed
|
||||||
|
except ImportError as e:
|
||||||
|
logger.error(
|
||||||
|
"deepspeed is not installed. please install deepspeed in your environment with following command. DS_BUILD_OPS=0 pip install deepspeed"
|
||||||
|
)
|
||||||
|
exit(1)
|
||||||
|
|
||||||
|
deepspeed_plugin = DeepSpeedPlugin(
|
||||||
|
zero_stage=args.zero_stage,
|
||||||
|
gradient_accumulation_steps=args.gradient_accumulation_steps,
|
||||||
|
gradient_clipping=args.max_grad_norm,
|
||||||
|
offload_optimizer_device=args.offload_optimizer_device,
|
||||||
|
offload_optimizer_nvme_path=args.offload_optimizer_nvme_path,
|
||||||
|
offload_param_device=args.offload_param_device,
|
||||||
|
offload_param_nvme_path=args.offload_param_nvme_path,
|
||||||
|
zero3_init_flag=args.zero3_init_flag,
|
||||||
|
zero3_save_16bit_model=args.zero3_save_16bit_model,
|
||||||
|
)
|
||||||
|
deepspeed_plugin.deepspeed_config["train_micro_batch_size_per_gpu"] = args.train_batch_size
|
||||||
|
deepspeed_plugin.deepspeed_config["train_batch_size"] = (
|
||||||
|
args.train_batch_size * args.gradient_accumulation_steps * int(os.environ["WORLD_SIZE"])
|
||||||
|
)
|
||||||
|
deepspeed_plugin.set_mixed_precision(args.mixed_precision)
|
||||||
|
if args.mixed_precision.lower() == "fp16":
|
||||||
|
deepspeed_plugin.deepspeed_config["fp16"]["initial_scale_power"] = 0 # preventing overflow.
|
||||||
|
if args.full_fp16 or args.fp16_master_weights_and_gradients:
|
||||||
|
if args.offload_optimizer_device == "cpu" and args.zero_stage == 2:
|
||||||
|
deepspeed_plugin.deepspeed_config["fp16"]["fp16_master_weights_and_grads"] = True
|
||||||
|
logger.info("[DeepSpeed] full fp16 enable.")
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
"[DeepSpeed]full fp16, fp16_master_weights_and_grads currently only supported using ZeRO-Offload with DeepSpeedCPUAdam on ZeRO-2 stage."
|
||||||
|
)
|
||||||
|
|
||||||
|
if args.offload_optimizer_device is not None:
|
||||||
|
logger.info("[DeepSpeed] start to manually build cpu_adam.")
|
||||||
|
deepspeed.ops.op_builder.CPUAdamBuilder().load()
|
||||||
|
logger.info("[DeepSpeed] building cpu_adam done.")
|
||||||
|
|
||||||
|
return deepspeed_plugin
|
||||||
|
|
||||||
|
|
||||||
|
# Accelerate library does not support multiple models for deepspeed. So, we need to wrap multiple models into a single model.
|
||||||
|
def prepare_deepspeed_model(args: argparse.Namespace, **models):
|
||||||
|
# remove None from models
|
||||||
|
models = {k: v for k, v in models.items() if v is not None}
|
||||||
|
|
||||||
|
class DeepSpeedWrapper(torch.nn.Module):
|
||||||
|
def __init__(self, **kw_models) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.models = torch.nn.ModuleDict()
|
||||||
|
|
||||||
|
for key, model in kw_models.items():
|
||||||
|
if isinstance(model, list):
|
||||||
|
model = torch.nn.ModuleList(model)
|
||||||
|
assert isinstance(
|
||||||
|
model, torch.nn.Module
|
||||||
|
), f"model must be an instance of torch.nn.Module, but got {key} is {type(model)}"
|
||||||
|
self.models.update(torch.nn.ModuleDict({key: model}))
|
||||||
|
|
||||||
|
def get_models(self):
|
||||||
|
return self.models
|
||||||
|
|
||||||
|
ds_model = DeepSpeedWrapper(**models)
|
||||||
|
return ds_model
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
import functools
|
||||||
|
import gc
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
try:
|
||||||
|
HAS_CUDA = torch.cuda.is_available()
|
||||||
|
except Exception:
|
||||||
|
HAS_CUDA = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
HAS_MPS = torch.backends.mps.is_available()
|
||||||
|
except Exception:
|
||||||
|
HAS_MPS = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
import intel_extension_for_pytorch as ipex # noqa
|
||||||
|
|
||||||
|
HAS_XPU = torch.xpu.is_available()
|
||||||
|
except Exception:
|
||||||
|
HAS_XPU = False
|
||||||
|
|
||||||
|
|
||||||
|
def clean_memory():
|
||||||
|
gc.collect()
|
||||||
|
if HAS_CUDA:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
if HAS_XPU:
|
||||||
|
torch.xpu.empty_cache()
|
||||||
|
if HAS_MPS:
|
||||||
|
torch.mps.empty_cache()
|
||||||
|
|
||||||
|
|
||||||
|
def clean_memory_on_device(device: torch.device):
|
||||||
|
r"""
|
||||||
|
Clean memory on the specified device, will be called from training scripts.
|
||||||
|
"""
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
# device may "cuda" or "cuda:0", so we need to check the type of device
|
||||||
|
if device.type == "cuda":
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
if device.type == "xpu":
|
||||||
|
torch.xpu.empty_cache()
|
||||||
|
if device.type == "mps":
|
||||||
|
torch.mps.empty_cache()
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=None)
|
||||||
|
def get_preferred_device() -> torch.device:
|
||||||
|
r"""
|
||||||
|
Do not call this function from training scripts. Use accelerator.device instead.
|
||||||
|
"""
|
||||||
|
if HAS_CUDA:
|
||||||
|
device = torch.device("cuda")
|
||||||
|
elif HAS_XPU:
|
||||||
|
device = torch.device("xpu")
|
||||||
|
elif HAS_MPS:
|
||||||
|
device = torch.device("mps")
|
||||||
|
else:
|
||||||
|
device = torch.device("cpu")
|
||||||
|
print(f"get_preferred_device() -> {device}")
|
||||||
|
return device
|
||||||
|
|
||||||
|
|
||||||
|
def init_ipex():
|
||||||
|
"""
|
||||||
|
Apply IPEX to CUDA hijacks using `library.ipex.ipex_init`.
|
||||||
|
|
||||||
|
This function should run right after importing torch and before doing anything else.
|
||||||
|
|
||||||
|
If IPEX is not available, this function does nothing.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if HAS_XPU:
|
||||||
|
from library.ipex import ipex_init
|
||||||
|
|
||||||
|
is_initialized, error_message = ipex_init()
|
||||||
|
if not is_initialized:
|
||||||
|
print("failed to initialize ipex:", error_message)
|
||||||
|
else:
|
||||||
|
return
|
||||||
|
except Exception as e:
|
||||||
|
print("failed to initialize ipex:", e)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,294 @@
|
|||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
import numpy as np
|
||||||
|
import toml
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from accelerate import Accelerator, PartialState
|
||||||
|
from transformers import CLIPTextModel
|
||||||
|
from tqdm import tqdm
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from . import flux_models, flux_utils, strategy_base
|
||||||
|
from .sd3_train_utils import load_prompts
|
||||||
|
from .device_utils import init_ipex, clean_memory_on_device
|
||||||
|
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def sample_images(
|
||||||
|
accelerator: Accelerator,
|
||||||
|
args: argparse.Namespace,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
flux,
|
||||||
|
ae,
|
||||||
|
text_encoders,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement=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
|
||||||
|
|
||||||
|
# unwrap unet and text_encoder(s)
|
||||||
|
flux = accelerator.unwrap_model(flux)
|
||||||
|
text_encoders = [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 = []
|
||||||
|
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)
|
||||||
|
|
||||||
|
save_dir = args.output_dir + "/sample"
|
||||||
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
|
||||||
|
# save random state to restore later
|
||||||
|
rng_state = torch.get_rng_state()
|
||||||
|
cuda_rng_state = None
|
||||||
|
try:
|
||||||
|
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
image_tensor_list = []
|
||||||
|
for prompt_dict in prompts:
|
||||||
|
image_tensor = sample_image_inference(
|
||||||
|
accelerator,
|
||||||
|
args,
|
||||||
|
flux,
|
||||||
|
text_encoders,
|
||||||
|
ae,
|
||||||
|
save_dir,
|
||||||
|
prompt_dict,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement,
|
||||||
|
)
|
||||||
|
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)
|
||||||
|
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
return torch.cat(image_tensor_list, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
def sample_image_inference(
|
||||||
|
accelerator: Accelerator,
|
||||||
|
args: argparse.Namespace,
|
||||||
|
flux: flux_models.Flux,
|
||||||
|
text_encoders: List[CLIPTextModel],
|
||||||
|
ae: flux_models.AutoEncoder,
|
||||||
|
save_dir,
|
||||||
|
prompt_dict,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement,
|
||||||
|
):
|
||||||
|
assert isinstance(prompt_dict, dict)
|
||||||
|
# negative_prompt = prompt_dict.get("negative_prompt")
|
||||||
|
sample_steps = prompt_dict.get("sample_steps", 20)
|
||||||
|
width = prompt_dict.get("width", 512)
|
||||||
|
height = prompt_dict.get("height", 512)
|
||||||
|
scale = prompt_dict.get("scale", 3.5)
|
||||||
|
seed = prompt_dict.get("seed")
|
||||||
|
# controlnet_image = prompt_dict.get("controlnet_image")
|
||||||
|
prompt: str = prompt_dict.get("prompt", "")
|
||||||
|
# sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler)
|
||||||
|
|
||||||
|
if prompt_replacement is not None:
|
||||||
|
prompt = prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||||
|
# if negative_prompt is not None:
|
||||||
|
# negative_prompt = negative_prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||||
|
|
||||||
|
if seed is not None:
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
torch.cuda.manual_seed(seed)
|
||||||
|
else:
|
||||||
|
# True random sample image generation
|
||||||
|
torch.seed()
|
||||||
|
torch.cuda.seed()
|
||||||
|
|
||||||
|
# if negative_prompt is None:
|
||||||
|
# negative_prompt = ""
|
||||||
|
|
||||||
|
height = max(64, height - height % 16) # round to divisible by 16
|
||||||
|
width = max(64, width - width % 16) # round to divisible by 16
|
||||||
|
logger.info(f"prompt: {prompt}")
|
||||||
|
# logger.info(f"negative_prompt: {negative_prompt}")
|
||||||
|
logger.info(f"height: {height}")
|
||||||
|
logger.info(f"width: {width}")
|
||||||
|
logger.info(f"sample_steps: {sample_steps}")
|
||||||
|
logger.info(f"scale: {scale}")
|
||||||
|
# logger.info(f"sample_sampler: {sampler_name}")
|
||||||
|
if seed is not None:
|
||||||
|
logger.info(f"seed: {seed}")
|
||||||
|
|
||||||
|
# encode prompts
|
||||||
|
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:
|
||||||
|
tokens_and_masks = tokenize_strategy.tokenize(prompt)
|
||||||
|
te_outputs = encoding_strategy.encode_tokens(tokenize_strategy, text_encoders, tokens_and_masks)
|
||||||
|
|
||||||
|
l_pooled, t5_out, txt_ids = te_outputs
|
||||||
|
|
||||||
|
# sample image
|
||||||
|
weight_dtype = ae.dtype # TOFO give dtype as argument
|
||||||
|
packed_latent_height = height // 16
|
||||||
|
packed_latent_width = width // 16
|
||||||
|
noise = torch.randn(
|
||||||
|
1,
|
||||||
|
packed_latent_height * packed_latent_width,
|
||||||
|
16 * 2 * 2,
|
||||||
|
device=accelerator.device,
|
||||||
|
dtype=weight_dtype,
|
||||||
|
generator=torch.Generator(device=accelerator.device).manual_seed(seed) if seed is not None else None,
|
||||||
|
)
|
||||||
|
timesteps = get_schedule(sample_steps, noise.shape[1], shift=True) # FLUX.1 dev -> shift=True
|
||||||
|
img_ids = flux_utils.prepare_img_ids(1, packed_latent_height, packed_latent_width).to(accelerator.device, weight_dtype)
|
||||||
|
|
||||||
|
with accelerator.autocast(), torch.no_grad():
|
||||||
|
x = denoise(flux, noise, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=scale)
|
||||||
|
|
||||||
|
x = x.float()
|
||||||
|
x = flux_utils.unpack_latents(x, packed_latent_height, packed_latent_width)
|
||||||
|
|
||||||
|
# latent to image
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
org_vae_device = ae.device # will be on cpu
|
||||||
|
ae.to(accelerator.device) # distributed_state.device is same as accelerator.device
|
||||||
|
with accelerator.autocast(), torch.no_grad():
|
||||||
|
x = ae.decode(x)
|
||||||
|
ae.to(org_vae_device)
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
|
||||||
|
x = x.clamp(-1, 1)
|
||||||
|
x = x.permute(0, 2, 3, 1)
|
||||||
|
image = Image.fromarray((127.5 * (x + 1.0)).float().cpu().numpy().astype(np.uint8)[0])
|
||||||
|
|
||||||
|
# adding accelerator.wait_for_everyone() here should sync up and ensure that sample images are saved in the same order as the original prompt list
|
||||||
|
# but adding 'enum' to the filename should be enough
|
||||||
|
|
||||||
|
ts_str = time.strftime("%Y%m%d%H%M%S", time.localtime())
|
||||||
|
num_suffix = f"e{epoch:06d}" if epoch is not None else f"{steps:06d}"
|
||||||
|
seed_suffix = "" if seed is None else f"_{seed}"
|
||||||
|
i: int = prompt_dict["enum"]
|
||||||
|
img_filename = f"{'' if args.output_name is None else args.output_name + '_'}{num_suffix}_{i:02d}_{ts_str}{seed_suffix}.png"
|
||||||
|
image.save(os.path.join(save_dir, img_filename))
|
||||||
|
return x
|
||||||
|
|
||||||
|
# wandb有効時のみログを送信
|
||||||
|
# try:
|
||||||
|
# wandb_tracker = accelerator.get_tracker("wandb")
|
||||||
|
# try:
|
||||||
|
# import wandb
|
||||||
|
# except ImportError: # 事前に一度確認するのでここはエラー出ないはず
|
||||||
|
# raise ImportError("No wandb / wandb がインストールされていないようです")
|
||||||
|
|
||||||
|
# wandb_tracker.log({f"sample_{i}": wandb.Image(image)})
|
||||||
|
# except: # wandb 無効時
|
||||||
|
# pass
|
||||||
|
|
||||||
|
|
||||||
|
def time_shift(mu: float, sigma: float, t: torch.Tensor):
|
||||||
|
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||||
|
|
||||||
|
|
||||||
|
def get_lin_function(x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15) -> Callable[[float], float]:
|
||||||
|
m = (y2 - y1) / (x2 - x1)
|
||||||
|
b = y1 - m * x1
|
||||||
|
return lambda x: m * x + b
|
||||||
|
|
||||||
|
|
||||||
|
def get_schedule(
|
||||||
|
num_steps: int,
|
||||||
|
image_seq_len: int,
|
||||||
|
base_shift: float = 0.5,
|
||||||
|
max_shift: float = 1.15,
|
||||||
|
shift: bool = True,
|
||||||
|
) -> list[float]:
|
||||||
|
# extra step for zero
|
||||||
|
timesteps = torch.linspace(1, 0, num_steps + 1)
|
||||||
|
|
||||||
|
# shifting the schedule to favor high timesteps for higher signal images
|
||||||
|
if shift:
|
||||||
|
# eastimate mu based on linear estimation between two points
|
||||||
|
mu = get_lin_function(y1=base_shift, y2=max_shift)(image_seq_len)
|
||||||
|
timesteps = time_shift(mu, 1.0, timesteps)
|
||||||
|
|
||||||
|
return timesteps.tolist()
|
||||||
|
|
||||||
|
|
||||||
|
def denoise(
|
||||||
|
model: flux_models.Flux,
|
||||||
|
img: torch.Tensor,
|
||||||
|
img_ids: torch.Tensor,
|
||||||
|
txt: torch.Tensor,
|
||||||
|
txt_ids: torch.Tensor,
|
||||||
|
vec: torch.Tensor,
|
||||||
|
timesteps: list[float],
|
||||||
|
guidance: float = 4.0,
|
||||||
|
):
|
||||||
|
# this is ignored for schnell
|
||||||
|
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
||||||
|
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)
|
||||||
|
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
||||||
|
|
||||||
|
img = img + (t_prev - t_curr) * pred
|
||||||
|
|
||||||
|
return img
|
||||||
@@ -0,0 +1,215 @@
|
|||||||
|
import json
|
||||||
|
from typing import Union
|
||||||
|
import einops
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
from accelerate import init_empty_weights
|
||||||
|
from transformers import CLIPTextModel, CLIPConfig, T5EncoderModel, T5Config
|
||||||
|
|
||||||
|
#from library import flux_models
|
||||||
|
from .flux_models import Flux, AutoEncoder, configs
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
MODEL_VERSION_FLUX_V1 = "flux1"
|
||||||
|
|
||||||
|
|
||||||
|
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}")
|
||||||
|
with torch.device("meta"):
|
||||||
|
model = Flux(configs[name].params).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))
|
||||||
|
info = model.load_state_dict(sd, strict=False, assign=True)
|
||||||
|
logger.info(f"Loaded Flux: {info}")
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def load_ae(name: str, ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> AutoEncoder:
|
||||||
|
logger.info("Building AutoEncoder")
|
||||||
|
with torch.device("meta"):
|
||||||
|
ae = AutoEncoder(configs[name].ae_params).to(dtype)
|
||||||
|
|
||||||
|
logger.info(f"Loading state dict from {ckpt_path}")
|
||||||
|
sd = load_file(ckpt_path, device=str(device))
|
||||||
|
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")
|
||||||
|
CLIPL_CONFIG = {
|
||||||
|
"_name_or_path": "clip-vit-large-patch14/",
|
||||||
|
"architectures": ["CLIPModel"],
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"logit_scale_init_value": 2.6592,
|
||||||
|
"model_type": "clip",
|
||||||
|
"projection_dim": 768,
|
||||||
|
# "text_config": {
|
||||||
|
"_name_or_path": "",
|
||||||
|
"add_cross_attention": False,
|
||||||
|
"architectures": None,
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bad_words_ids": None,
|
||||||
|
"bos_token_id": 0,
|
||||||
|
"chunk_size_feed_forward": 0,
|
||||||
|
"cross_attention_hidden_size": None,
|
||||||
|
"decoder_start_token_id": None,
|
||||||
|
"diversity_penalty": 0.0,
|
||||||
|
"do_sample": False,
|
||||||
|
"dropout": 0.0,
|
||||||
|
"early_stopping": False,
|
||||||
|
"encoder_no_repeat_ngram_size": 0,
|
||||||
|
"eos_token_id": 2,
|
||||||
|
"finetuning_task": None,
|
||||||
|
"forced_bos_token_id": None,
|
||||||
|
"forced_eos_token_id": None,
|
||||||
|
"hidden_act": "quick_gelu",
|
||||||
|
"hidden_size": 768,
|
||||||
|
"id2label": {"0": "LABEL_0", "1": "LABEL_1"},
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 3072,
|
||||||
|
"is_decoder": False,
|
||||||
|
"is_encoder_decoder": False,
|
||||||
|
"label2id": {"LABEL_0": 0, "LABEL_1": 1},
|
||||||
|
"layer_norm_eps": 1e-05,
|
||||||
|
"length_penalty": 1.0,
|
||||||
|
"max_length": 20,
|
||||||
|
"max_position_embeddings": 77,
|
||||||
|
"min_length": 0,
|
||||||
|
"model_type": "clip_text_model",
|
||||||
|
"no_repeat_ngram_size": 0,
|
||||||
|
"num_attention_heads": 12,
|
||||||
|
"num_beam_groups": 1,
|
||||||
|
"num_beams": 1,
|
||||||
|
"num_hidden_layers": 12,
|
||||||
|
"num_return_sequences": 1,
|
||||||
|
"output_attentions": False,
|
||||||
|
"output_hidden_states": False,
|
||||||
|
"output_scores": False,
|
||||||
|
"pad_token_id": 1,
|
||||||
|
"prefix": None,
|
||||||
|
"problem_type": None,
|
||||||
|
"projection_dim": 768,
|
||||||
|
"pruned_heads": {},
|
||||||
|
"remove_invalid_values": False,
|
||||||
|
"repetition_penalty": 1.0,
|
||||||
|
"return_dict": True,
|
||||||
|
"return_dict_in_generate": False,
|
||||||
|
"sep_token_id": None,
|
||||||
|
"task_specific_params": None,
|
||||||
|
"temperature": 1.0,
|
||||||
|
"tie_encoder_decoder": False,
|
||||||
|
"tie_word_embeddings": True,
|
||||||
|
"tokenizer_class": None,
|
||||||
|
"top_k": 50,
|
||||||
|
"top_p": 1.0,
|
||||||
|
"torch_dtype": None,
|
||||||
|
"torchscript": False,
|
||||||
|
"transformers_version": "4.16.0.dev0",
|
||||||
|
"use_bfloat16": False,
|
||||||
|
"vocab_size": 49408,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_size": 1280,
|
||||||
|
"intermediate_size": 5120,
|
||||||
|
"num_attention_heads": 20,
|
||||||
|
"num_hidden_layers": 32,
|
||||||
|
# },
|
||||||
|
# "text_config_dict": {
|
||||||
|
"hidden_size": 768,
|
||||||
|
"intermediate_size": 3072,
|
||||||
|
"num_attention_heads": 12,
|
||||||
|
"num_hidden_layers": 12,
|
||||||
|
"projection_dim": 768,
|
||||||
|
# },
|
||||||
|
# "torch_dtype": "float32",
|
||||||
|
# "transformers_version": None,
|
||||||
|
}
|
||||||
|
config = CLIPConfig(**CLIPL_CONFIG)
|
||||||
|
with init_empty_weights():
|
||||||
|
clip = CLIPTextModel._from_config(config)
|
||||||
|
|
||||||
|
logger.info(f"Loading state dict from {ckpt_path}")
|
||||||
|
sd = load_file(ckpt_path, device=str(device))
|
||||||
|
info = clip.load_state_dict(sd, strict=False, assign=True)
|
||||||
|
logger.info(f"Loaded CLIP: {info}")
|
||||||
|
return clip
|
||||||
|
|
||||||
|
|
||||||
|
def load_t5xxl(ckpt_path: str, dtype: torch.dtype, device: Union[str, torch.device]) -> T5EncoderModel:
|
||||||
|
T5_CONFIG_JSON = """
|
||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"T5EncoderModel"
|
||||||
|
],
|
||||||
|
"classifier_dropout": 0.0,
|
||||||
|
"d_ff": 10240,
|
||||||
|
"d_kv": 64,
|
||||||
|
"d_model": 4096,
|
||||||
|
"decoder_start_token_id": 0,
|
||||||
|
"dense_act_fn": "gelu_new",
|
||||||
|
"dropout_rate": 0.1,
|
||||||
|
"eos_token_id": 1,
|
||||||
|
"feed_forward_proj": "gated-gelu",
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"is_encoder_decoder": true,
|
||||||
|
"is_gated_act": true,
|
||||||
|
"layer_norm_epsilon": 1e-06,
|
||||||
|
"model_type": "t5",
|
||||||
|
"num_decoder_layers": 24,
|
||||||
|
"num_heads": 64,
|
||||||
|
"num_layers": 24,
|
||||||
|
"output_past": true,
|
||||||
|
"pad_token_id": 0,
|
||||||
|
"relative_attention_max_distance": 128,
|
||||||
|
"relative_attention_num_buckets": 32,
|
||||||
|
"tie_word_embeddings": false,
|
||||||
|
"torch_dtype": "float16",
|
||||||
|
"transformers_version": "4.41.2",
|
||||||
|
"use_cache": true,
|
||||||
|
"vocab_size": 32128
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
config = json.loads(T5_CONFIG_JSON)
|
||||||
|
config = T5Config(**config)
|
||||||
|
with init_empty_weights():
|
||||||
|
t5xxl = T5EncoderModel._from_config(config)
|
||||||
|
|
||||||
|
logger.info(f"Loading state dict from {ckpt_path}")
|
||||||
|
sd = load_file(ckpt_path, device=str(device))
|
||||||
|
info = t5xxl.load_state_dict(sd, strict=False, assign=True)
|
||||||
|
logger.info(f"Loaded T5xxl: {info}")
|
||||||
|
return t5xxl
|
||||||
|
|
||||||
|
|
||||||
|
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]
|
||||||
|
img_ids[..., 2] = img_ids[..., 2] + torch.arange(packed_latent_width)[None, :]
|
||||||
|
img_ids = einops.repeat(img_ids, "h w c -> b (h w) c", b=batch_size)
|
||||||
|
return img_ids
|
||||||
|
|
||||||
|
|
||||||
|
def unpack_latents(x: torch.Tensor, packed_latent_height: int, packed_latent_width: int) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
x: [b (h w) (c ph pw)] -> [b c (h ph) (w pw)], ph=2, pw=2
|
||||||
|
"""
|
||||||
|
x = einops.rearrange(x, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=packed_latent_height, w=packed_latent_width, ph=2, pw=2)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def pack_latents(x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
x: [b c (h ph) (w pw)] -> [b (h w) (c ph pw)], ph=2, pw=2
|
||||||
|
"""
|
||||||
|
x = einops.rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
|
||||||
|
return x
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
from typing import Union, BinaryIO
|
||||||
|
from huggingface_hub import HfApi
|
||||||
|
from pathlib import Path
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
from .utils import fire_in_thread
|
||||||
|
from .utils import setup_logging
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def exists_repo(repo_id: str, repo_type: str, revision: str = "main", token: str = None):
|
||||||
|
api = HfApi(
|
||||||
|
token=token,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
api.repo_info(repo_id=repo_id, revision=revision, repo_type=repo_type)
|
||||||
|
return True
|
||||||
|
except:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def upload(
|
||||||
|
args: argparse.Namespace,
|
||||||
|
src: Union[str, Path, bytes, BinaryIO],
|
||||||
|
dest_suffix: str = "",
|
||||||
|
force_sync_upload: bool = False,
|
||||||
|
):
|
||||||
|
repo_id = args.huggingface_repo_id
|
||||||
|
repo_type = args.huggingface_repo_type
|
||||||
|
token = args.huggingface_token
|
||||||
|
path_in_repo = args.huggingface_path_in_repo + dest_suffix if args.huggingface_path_in_repo is not None else None
|
||||||
|
private = args.huggingface_repo_visibility is None or args.huggingface_repo_visibility != "public"
|
||||||
|
api = HfApi(token=token)
|
||||||
|
if not exists_repo(repo_id=repo_id, repo_type=repo_type, token=token):
|
||||||
|
try:
|
||||||
|
api.create_repo(repo_id=repo_id, repo_type=repo_type, private=private)
|
||||||
|
except Exception as e: # とりあえずRepositoryNotFoundErrorは確認したが他にあると困るので
|
||||||
|
logger.error("===========================================")
|
||||||
|
logger.error(f"failed to create HuggingFace repo / HuggingFaceのリポジトリの作成に失敗しました : {e}")
|
||||||
|
logger.error("===========================================")
|
||||||
|
|
||||||
|
is_folder = (type(src) == str and os.path.isdir(src)) or (isinstance(src, Path) and src.is_dir())
|
||||||
|
|
||||||
|
def uploader():
|
||||||
|
try:
|
||||||
|
if is_folder:
|
||||||
|
api.upload_folder(
|
||||||
|
repo_id=repo_id,
|
||||||
|
repo_type=repo_type,
|
||||||
|
folder_path=src,
|
||||||
|
path_in_repo=path_in_repo,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
api.upload_file(
|
||||||
|
repo_id=repo_id,
|
||||||
|
repo_type=repo_type,
|
||||||
|
path_or_fileobj=src,
|
||||||
|
path_in_repo=path_in_repo,
|
||||||
|
)
|
||||||
|
except Exception as e: # RuntimeErrorを確認済みだが他にあると困るので
|
||||||
|
logger.error("===========================================")
|
||||||
|
logger.error(f"failed to upload to HuggingFace / HuggingFaceへのアップロードに失敗しました : {e}")
|
||||||
|
logger.error("===========================================")
|
||||||
|
|
||||||
|
if args.async_upload and not force_sync_upload:
|
||||||
|
fire_in_thread(uploader)
|
||||||
|
else:
|
||||||
|
uploader()
|
||||||
|
|
||||||
|
|
||||||
|
def list_dir(
|
||||||
|
repo_id: str,
|
||||||
|
subfolder: str,
|
||||||
|
repo_type: str,
|
||||||
|
revision: str = "main",
|
||||||
|
token: str = None,
|
||||||
|
):
|
||||||
|
api = HfApi(
|
||||||
|
token=token,
|
||||||
|
)
|
||||||
|
repo_info = api.repo_info(repo_id=repo_id, revision=revision, repo_type=repo_type)
|
||||||
|
file_list = [file for file in repo_info.siblings if file.rfilename.startswith(subfolder)]
|
||||||
|
return file_list
|
||||||
@@ -0,0 +1,223 @@
|
|||||||
|
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
|
||||||
@@ -0,0 +1,180 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import contextlib
|
||||||
|
import torch
|
||||||
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||||
|
from .hijacks import ipex_hijacks
|
||||||
|
|
||||||
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
||||||
|
|
||||||
|
def ipex_init(): # pylint: disable=too-many-statements
|
||||||
|
try:
|
||||||
|
if hasattr(torch, "cuda") and hasattr(torch.cuda, "is_xpu_hijacked") and torch.cuda.is_xpu_hijacked:
|
||||||
|
return True, "Skipping IPEX hijack"
|
||||||
|
else:
|
||||||
|
# Replace cuda with xpu:
|
||||||
|
torch.cuda.current_device = torch.xpu.current_device
|
||||||
|
torch.cuda.current_stream = torch.xpu.current_stream
|
||||||
|
torch.cuda.device = torch.xpu.device
|
||||||
|
torch.cuda.device_count = torch.xpu.device_count
|
||||||
|
torch.cuda.device_of = torch.xpu.device_of
|
||||||
|
torch.cuda.get_device_name = torch.xpu.get_device_name
|
||||||
|
torch.cuda.get_device_properties = torch.xpu.get_device_properties
|
||||||
|
torch.cuda.init = torch.xpu.init
|
||||||
|
torch.cuda.is_available = torch.xpu.is_available
|
||||||
|
torch.cuda.is_initialized = torch.xpu.is_initialized
|
||||||
|
torch.cuda.is_current_stream_capturing = lambda: False
|
||||||
|
torch.cuda.set_device = torch.xpu.set_device
|
||||||
|
torch.cuda.stream = torch.xpu.stream
|
||||||
|
torch.cuda.synchronize = torch.xpu.synchronize
|
||||||
|
torch.cuda.Event = torch.xpu.Event
|
||||||
|
torch.cuda.Stream = torch.xpu.Stream
|
||||||
|
torch.cuda.FloatTensor = torch.xpu.FloatTensor
|
||||||
|
torch.Tensor.cuda = torch.Tensor.xpu
|
||||||
|
torch.Tensor.is_cuda = torch.Tensor.is_xpu
|
||||||
|
torch.nn.Module.cuda = torch.nn.Module.xpu
|
||||||
|
torch.UntypedStorage.cuda = torch.UntypedStorage.xpu
|
||||||
|
torch.cuda._initialization_lock = torch.xpu.lazy_init._initialization_lock
|
||||||
|
torch.cuda._initialized = torch.xpu.lazy_init._initialized
|
||||||
|
torch.cuda._lazy_seed_tracker = torch.xpu.lazy_init._lazy_seed_tracker
|
||||||
|
torch.cuda._queued_calls = torch.xpu.lazy_init._queued_calls
|
||||||
|
torch.cuda._tls = torch.xpu.lazy_init._tls
|
||||||
|
torch.cuda.threading = torch.xpu.lazy_init.threading
|
||||||
|
torch.cuda.traceback = torch.xpu.lazy_init.traceback
|
||||||
|
torch.cuda.Optional = torch.xpu.Optional
|
||||||
|
torch.cuda.__cached__ = torch.xpu.__cached__
|
||||||
|
torch.cuda.__loader__ = torch.xpu.__loader__
|
||||||
|
torch.cuda.ComplexFloatStorage = torch.xpu.ComplexFloatStorage
|
||||||
|
torch.cuda.Tuple = torch.xpu.Tuple
|
||||||
|
torch.cuda.streams = torch.xpu.streams
|
||||||
|
torch.cuda._lazy_new = torch.xpu._lazy_new
|
||||||
|
torch.cuda.FloatStorage = torch.xpu.FloatStorage
|
||||||
|
torch.cuda.Any = torch.xpu.Any
|
||||||
|
torch.cuda.__doc__ = torch.xpu.__doc__
|
||||||
|
torch.cuda.default_generators = torch.xpu.default_generators
|
||||||
|
torch.cuda.HalfTensor = torch.xpu.HalfTensor
|
||||||
|
torch.cuda._get_device_index = torch.xpu._get_device_index
|
||||||
|
torch.cuda.__path__ = torch.xpu.__path__
|
||||||
|
torch.cuda.Device = torch.xpu.Device
|
||||||
|
torch.cuda.IntTensor = torch.xpu.IntTensor
|
||||||
|
torch.cuda.ByteStorage = torch.xpu.ByteStorage
|
||||||
|
torch.cuda.set_stream = torch.xpu.set_stream
|
||||||
|
torch.cuda.BoolStorage = torch.xpu.BoolStorage
|
||||||
|
torch.cuda.os = torch.xpu.os
|
||||||
|
torch.cuda.torch = torch.xpu.torch
|
||||||
|
torch.cuda.BFloat16Storage = torch.xpu.BFloat16Storage
|
||||||
|
torch.cuda.Union = torch.xpu.Union
|
||||||
|
torch.cuda.DoubleTensor = torch.xpu.DoubleTensor
|
||||||
|
torch.cuda.ShortTensor = torch.xpu.ShortTensor
|
||||||
|
torch.cuda.LongTensor = torch.xpu.LongTensor
|
||||||
|
torch.cuda.IntStorage = torch.xpu.IntStorage
|
||||||
|
torch.cuda.LongStorage = torch.xpu.LongStorage
|
||||||
|
torch.cuda.__annotations__ = torch.xpu.__annotations__
|
||||||
|
torch.cuda.__package__ = torch.xpu.__package__
|
||||||
|
torch.cuda.__builtins__ = torch.xpu.__builtins__
|
||||||
|
torch.cuda.CharTensor = torch.xpu.CharTensor
|
||||||
|
torch.cuda.List = torch.xpu.List
|
||||||
|
torch.cuda._lazy_init = torch.xpu._lazy_init
|
||||||
|
torch.cuda.BFloat16Tensor = torch.xpu.BFloat16Tensor
|
||||||
|
torch.cuda.DoubleStorage = torch.xpu.DoubleStorage
|
||||||
|
torch.cuda.ByteTensor = torch.xpu.ByteTensor
|
||||||
|
torch.cuda.StreamContext = torch.xpu.StreamContext
|
||||||
|
torch.cuda.ComplexDoubleStorage = torch.xpu.ComplexDoubleStorage
|
||||||
|
torch.cuda.ShortStorage = torch.xpu.ShortStorage
|
||||||
|
torch.cuda._lazy_call = torch.xpu._lazy_call
|
||||||
|
torch.cuda.HalfStorage = torch.xpu.HalfStorage
|
||||||
|
torch.cuda.random = torch.xpu.random
|
||||||
|
torch.cuda._device = torch.xpu._device
|
||||||
|
torch.cuda.classproperty = torch.xpu.classproperty
|
||||||
|
torch.cuda.__name__ = torch.xpu.__name__
|
||||||
|
torch.cuda._device_t = torch.xpu._device_t
|
||||||
|
torch.cuda.warnings = torch.xpu.warnings
|
||||||
|
torch.cuda.__spec__ = torch.xpu.__spec__
|
||||||
|
torch.cuda.BoolTensor = torch.xpu.BoolTensor
|
||||||
|
torch.cuda.CharStorage = torch.xpu.CharStorage
|
||||||
|
torch.cuda.__file__ = torch.xpu.__file__
|
||||||
|
torch.cuda._is_in_bad_fork = torch.xpu.lazy_init._is_in_bad_fork
|
||||||
|
# torch.cuda.is_current_stream_capturing = torch.xpu.is_current_stream_capturing
|
||||||
|
|
||||||
|
# Memory:
|
||||||
|
torch.cuda.memory = torch.xpu.memory
|
||||||
|
if 'linux' in sys.platform and "WSL2" in os.popen("uname -a").read():
|
||||||
|
torch.xpu.empty_cache = lambda: None
|
||||||
|
torch.cuda.empty_cache = torch.xpu.empty_cache
|
||||||
|
torch.cuda.memory_stats = torch.xpu.memory_stats
|
||||||
|
torch.cuda.memory_summary = torch.xpu.memory_summary
|
||||||
|
torch.cuda.memory_snapshot = torch.xpu.memory_snapshot
|
||||||
|
torch.cuda.memory_allocated = torch.xpu.memory_allocated
|
||||||
|
torch.cuda.max_memory_allocated = torch.xpu.max_memory_allocated
|
||||||
|
torch.cuda.memory_reserved = torch.xpu.memory_reserved
|
||||||
|
torch.cuda.memory_cached = torch.xpu.memory_reserved
|
||||||
|
torch.cuda.max_memory_reserved = torch.xpu.max_memory_reserved
|
||||||
|
torch.cuda.max_memory_cached = torch.xpu.max_memory_reserved
|
||||||
|
torch.cuda.reset_peak_memory_stats = torch.xpu.reset_peak_memory_stats
|
||||||
|
torch.cuda.reset_max_memory_cached = torch.xpu.reset_peak_memory_stats
|
||||||
|
torch.cuda.reset_max_memory_allocated = torch.xpu.reset_peak_memory_stats
|
||||||
|
torch.cuda.memory_stats_as_nested_dict = torch.xpu.memory_stats_as_nested_dict
|
||||||
|
torch.cuda.reset_accumulated_memory_stats = torch.xpu.reset_accumulated_memory_stats
|
||||||
|
|
||||||
|
# RNG:
|
||||||
|
torch.cuda.get_rng_state = torch.xpu.get_rng_state
|
||||||
|
torch.cuda.get_rng_state_all = torch.xpu.get_rng_state_all
|
||||||
|
torch.cuda.set_rng_state = torch.xpu.set_rng_state
|
||||||
|
torch.cuda.set_rng_state_all = torch.xpu.set_rng_state_all
|
||||||
|
torch.cuda.manual_seed = torch.xpu.manual_seed
|
||||||
|
torch.cuda.manual_seed_all = torch.xpu.manual_seed_all
|
||||||
|
torch.cuda.seed = torch.xpu.seed
|
||||||
|
torch.cuda.seed_all = torch.xpu.seed_all
|
||||||
|
torch.cuda.initial_seed = torch.xpu.initial_seed
|
||||||
|
|
||||||
|
# AMP:
|
||||||
|
torch.cuda.amp = torch.xpu.amp
|
||||||
|
torch.is_autocast_enabled = torch.xpu.is_autocast_xpu_enabled
|
||||||
|
torch.get_autocast_gpu_dtype = torch.xpu.get_autocast_xpu_dtype
|
||||||
|
|
||||||
|
if not hasattr(torch.cuda.amp, "common"):
|
||||||
|
torch.cuda.amp.common = contextlib.nullcontext()
|
||||||
|
torch.cuda.amp.common.amp_definitely_not_available = lambda: False
|
||||||
|
|
||||||
|
try:
|
||||||
|
torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler
|
||||||
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
|
try:
|
||||||
|
from .gradscaler import gradscaler_init # pylint: disable=import-outside-toplevel, import-error
|
||||||
|
gradscaler_init()
|
||||||
|
torch.cuda.amp.GradScaler = torch.xpu.amp.GradScaler
|
||||||
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
|
torch.cuda.amp.GradScaler = ipex.cpu.autocast._grad_scaler.GradScaler
|
||||||
|
|
||||||
|
# C
|
||||||
|
torch._C._cuda_getCurrentRawStream = ipex._C._getCurrentStream
|
||||||
|
ipex._C._DeviceProperties.multi_processor_count = ipex._C._DeviceProperties.gpu_subslice_count
|
||||||
|
ipex._C._DeviceProperties.major = 2024
|
||||||
|
ipex._C._DeviceProperties.minor = 0
|
||||||
|
|
||||||
|
# Fix functions with ipex:
|
||||||
|
torch.cuda.mem_get_info = lambda device=None: [(torch.xpu.get_device_properties(device).total_memory - torch.xpu.memory_reserved(device)), torch.xpu.get_device_properties(device).total_memory]
|
||||||
|
torch._utils._get_available_device_type = lambda: "xpu"
|
||||||
|
torch.has_cuda = True
|
||||||
|
torch.cuda.has_half = True
|
||||||
|
torch.cuda.is_bf16_supported = lambda *args, **kwargs: True
|
||||||
|
torch.cuda.is_fp16_supported = lambda *args, **kwargs: True
|
||||||
|
torch.backends.cuda.is_built = lambda *args, **kwargs: True
|
||||||
|
torch.version.cuda = "12.1"
|
||||||
|
torch.cuda.get_device_capability = lambda *args, **kwargs: [12,1]
|
||||||
|
torch.cuda.get_device_properties.major = 12
|
||||||
|
torch.cuda.get_device_properties.minor = 1
|
||||||
|
torch.cuda.ipc_collect = lambda *args, **kwargs: None
|
||||||
|
torch.cuda.utilization = lambda *args, **kwargs: 0
|
||||||
|
|
||||||
|
ipex_hijacks()
|
||||||
|
if not torch.xpu.has_fp64_dtype() or os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None) is not None:
|
||||||
|
try:
|
||||||
|
from .diffusers import ipex_diffusers
|
||||||
|
ipex_diffusers()
|
||||||
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
|
pass
|
||||||
|
torch.cuda.is_xpu_hijacked = True
|
||||||
|
except Exception as e:
|
||||||
|
return False, e
|
||||||
|
return True, None
|
||||||
@@ -0,0 +1,177 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||||
|
from functools import cache
|
||||||
|
|
||||||
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
||||||
|
|
||||||
|
# ARC GPUs can't allocate more than 4GB to a single block so we slice the attention layers
|
||||||
|
|
||||||
|
sdpa_slice_trigger_rate = float(os.environ.get('IPEX_SDPA_SLICE_TRIGGER_RATE', 4))
|
||||||
|
attention_slice_rate = float(os.environ.get('IPEX_ATTENTION_SLICE_RATE', 4))
|
||||||
|
|
||||||
|
# Find something divisible with the input_tokens
|
||||||
|
@cache
|
||||||
|
def find_slice_size(slice_size, slice_block_size):
|
||||||
|
while (slice_size * slice_block_size) > attention_slice_rate:
|
||||||
|
slice_size = slice_size // 2
|
||||||
|
if slice_size <= 1:
|
||||||
|
slice_size = 1
|
||||||
|
break
|
||||||
|
return slice_size
|
||||||
|
|
||||||
|
# Find slice sizes for SDPA
|
||||||
|
@cache
|
||||||
|
def find_sdpa_slice_sizes(query_shape, query_element_size):
|
||||||
|
if len(query_shape) == 3:
|
||||||
|
batch_size_attention, query_tokens, shape_three = query_shape
|
||||||
|
shape_four = 1
|
||||||
|
else:
|
||||||
|
batch_size_attention, query_tokens, shape_three, shape_four = query_shape
|
||||||
|
|
||||||
|
slice_block_size = query_tokens * shape_three * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
block_size = batch_size_attention * slice_block_size
|
||||||
|
|
||||||
|
split_slice_size = batch_size_attention
|
||||||
|
split_2_slice_size = query_tokens
|
||||||
|
split_3_slice_size = shape_three
|
||||||
|
|
||||||
|
do_split = False
|
||||||
|
do_split_2 = False
|
||||||
|
do_split_3 = False
|
||||||
|
|
||||||
|
if block_size > sdpa_slice_trigger_rate:
|
||||||
|
do_split = True
|
||||||
|
split_slice_size = find_slice_size(split_slice_size, slice_block_size)
|
||||||
|
if split_slice_size * slice_block_size > attention_slice_rate:
|
||||||
|
slice_2_block_size = split_slice_size * shape_three * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
do_split_2 = True
|
||||||
|
split_2_slice_size = find_slice_size(split_2_slice_size, slice_2_block_size)
|
||||||
|
if split_2_slice_size * slice_2_block_size > attention_slice_rate:
|
||||||
|
slice_3_block_size = split_slice_size * split_2_slice_size * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
do_split_3 = True
|
||||||
|
split_3_slice_size = find_slice_size(split_3_slice_size, slice_3_block_size)
|
||||||
|
|
||||||
|
return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size
|
||||||
|
|
||||||
|
# Find slice sizes for BMM
|
||||||
|
@cache
|
||||||
|
def find_bmm_slice_sizes(input_shape, input_element_size, mat2_shape):
|
||||||
|
batch_size_attention, input_tokens, mat2_atten_shape = input_shape[0], input_shape[1], mat2_shape[2]
|
||||||
|
slice_block_size = input_tokens * mat2_atten_shape / 1024 / 1024 * input_element_size
|
||||||
|
block_size = batch_size_attention * slice_block_size
|
||||||
|
|
||||||
|
split_slice_size = batch_size_attention
|
||||||
|
split_2_slice_size = input_tokens
|
||||||
|
split_3_slice_size = mat2_atten_shape
|
||||||
|
|
||||||
|
do_split = False
|
||||||
|
do_split_2 = False
|
||||||
|
do_split_3 = False
|
||||||
|
|
||||||
|
if block_size > attention_slice_rate:
|
||||||
|
do_split = True
|
||||||
|
split_slice_size = find_slice_size(split_slice_size, slice_block_size)
|
||||||
|
if split_slice_size * slice_block_size > attention_slice_rate:
|
||||||
|
slice_2_block_size = split_slice_size * mat2_atten_shape / 1024 / 1024 * input_element_size
|
||||||
|
do_split_2 = True
|
||||||
|
split_2_slice_size = find_slice_size(split_2_slice_size, slice_2_block_size)
|
||||||
|
if split_2_slice_size * slice_2_block_size > attention_slice_rate:
|
||||||
|
slice_3_block_size = split_slice_size * split_2_slice_size / 1024 / 1024 * input_element_size
|
||||||
|
do_split_3 = True
|
||||||
|
split_3_slice_size = find_slice_size(split_3_slice_size, slice_3_block_size)
|
||||||
|
|
||||||
|
return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size
|
||||||
|
|
||||||
|
|
||||||
|
original_torch_bmm = torch.bmm
|
||||||
|
def torch_bmm_32_bit(input, mat2, *, out=None):
|
||||||
|
if input.device.type != "xpu":
|
||||||
|
return original_torch_bmm(input, mat2, out=out)
|
||||||
|
do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_bmm_slice_sizes(input.shape, input.element_size(), mat2.shape)
|
||||||
|
|
||||||
|
# Slice BMM
|
||||||
|
if do_split:
|
||||||
|
batch_size_attention, input_tokens, mat2_atten_shape = input.shape[0], input.shape[1], mat2.shape[2]
|
||||||
|
hidden_states = torch.zeros(input.shape[0], input.shape[1], mat2.shape[2], device=input.device, dtype=input.dtype)
|
||||||
|
for i in range(batch_size_attention // split_slice_size):
|
||||||
|
start_idx = i * split_slice_size
|
||||||
|
end_idx = (i + 1) * split_slice_size
|
||||||
|
if do_split_2:
|
||||||
|
for i2 in range(input_tokens // split_2_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_2 = i2 * split_2_slice_size
|
||||||
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
||||||
|
if do_split_3:
|
||||||
|
for i3 in range(mat2_atten_shape // split_3_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_3 = i3 * split_3_slice_size
|
||||||
|
end_idx_3 = (i3 + 1) * split_3_slice_size
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] = original_torch_bmm(
|
||||||
|
input[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3],
|
||||||
|
mat2[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3],
|
||||||
|
out=out
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = original_torch_bmm(
|
||||||
|
input[start_idx:end_idx, start_idx_2:end_idx_2],
|
||||||
|
mat2[start_idx:end_idx, start_idx_2:end_idx_2],
|
||||||
|
out=out
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states[start_idx:end_idx] = original_torch_bmm(
|
||||||
|
input[start_idx:end_idx],
|
||||||
|
mat2[start_idx:end_idx],
|
||||||
|
out=out
|
||||||
|
)
|
||||||
|
torch.xpu.synchronize(input.device)
|
||||||
|
else:
|
||||||
|
return original_torch_bmm(input, mat2, out=out)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention
|
||||||
|
def scaled_dot_product_attention_32_bit(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, **kwargs):
|
||||||
|
if query.device.type != "xpu":
|
||||||
|
return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs)
|
||||||
|
do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_sdpa_slice_sizes(query.shape, query.element_size())
|
||||||
|
|
||||||
|
# Slice SDPA
|
||||||
|
if do_split:
|
||||||
|
batch_size_attention, query_tokens, shape_three = query.shape[0], query.shape[1], query.shape[2]
|
||||||
|
hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype)
|
||||||
|
for i in range(batch_size_attention // split_slice_size):
|
||||||
|
start_idx = i * split_slice_size
|
||||||
|
end_idx = (i + 1) * split_slice_size
|
||||||
|
if do_split_2:
|
||||||
|
for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_2 = i2 * split_2_slice_size
|
||||||
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
||||||
|
if do_split_3:
|
||||||
|
for i3 in range(shape_three // split_3_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_3 = i3 * split_3_slice_size
|
||||||
|
end_idx_3 = (i3 + 1) * split_3_slice_size
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] = original_scaled_dot_product_attention(
|
||||||
|
query[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3],
|
||||||
|
key[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3],
|
||||||
|
value[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3],
|
||||||
|
attn_mask=attn_mask[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] if attn_mask is not None else attn_mask,
|
||||||
|
dropout_p=dropout_p, is_causal=is_causal, **kwargs
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = original_scaled_dot_product_attention(
|
||||||
|
query[start_idx:end_idx, start_idx_2:end_idx_2],
|
||||||
|
key[start_idx:end_idx, start_idx_2:end_idx_2],
|
||||||
|
value[start_idx:end_idx, start_idx_2:end_idx_2],
|
||||||
|
attn_mask=attn_mask[start_idx:end_idx, start_idx_2:end_idx_2] if attn_mask is not None else attn_mask,
|
||||||
|
dropout_p=dropout_p, is_causal=is_causal, **kwargs
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
hidden_states[start_idx:end_idx] = original_scaled_dot_product_attention(
|
||||||
|
query[start_idx:end_idx],
|
||||||
|
key[start_idx:end_idx],
|
||||||
|
value[start_idx:end_idx],
|
||||||
|
attn_mask=attn_mask[start_idx:end_idx] if attn_mask is not None else attn_mask,
|
||||||
|
dropout_p=dropout_p, is_causal=is_causal, **kwargs
|
||||||
|
)
|
||||||
|
torch.xpu.synchronize(query.device)
|
||||||
|
else:
|
||||||
|
return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs)
|
||||||
|
return hidden_states
|
||||||
@@ -0,0 +1,312 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||||
|
import diffusers #0.24.0 # pylint: disable=import-error
|
||||||
|
from diffusers.models.attention_processor import Attention
|
||||||
|
from diffusers.utils import USE_PEFT_BACKEND
|
||||||
|
from functools import cache
|
||||||
|
|
||||||
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
||||||
|
|
||||||
|
attention_slice_rate = float(os.environ.get('IPEX_ATTENTION_SLICE_RATE', 4))
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def find_slice_size(slice_size, slice_block_size):
|
||||||
|
while (slice_size * slice_block_size) > attention_slice_rate:
|
||||||
|
slice_size = slice_size // 2
|
||||||
|
if slice_size <= 1:
|
||||||
|
slice_size = 1
|
||||||
|
break
|
||||||
|
return slice_size
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def find_attention_slice_sizes(query_shape, query_element_size, query_device_type, slice_size=None):
|
||||||
|
if len(query_shape) == 3:
|
||||||
|
batch_size_attention, query_tokens, shape_three = query_shape
|
||||||
|
shape_four = 1
|
||||||
|
else:
|
||||||
|
batch_size_attention, query_tokens, shape_three, shape_four = query_shape
|
||||||
|
if slice_size is not None:
|
||||||
|
batch_size_attention = slice_size
|
||||||
|
|
||||||
|
slice_block_size = query_tokens * shape_three * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
block_size = batch_size_attention * slice_block_size
|
||||||
|
|
||||||
|
split_slice_size = batch_size_attention
|
||||||
|
split_2_slice_size = query_tokens
|
||||||
|
split_3_slice_size = shape_three
|
||||||
|
|
||||||
|
do_split = False
|
||||||
|
do_split_2 = False
|
||||||
|
do_split_3 = False
|
||||||
|
|
||||||
|
if query_device_type != "xpu":
|
||||||
|
return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size
|
||||||
|
|
||||||
|
if block_size > attention_slice_rate:
|
||||||
|
do_split = True
|
||||||
|
split_slice_size = find_slice_size(split_slice_size, slice_block_size)
|
||||||
|
if split_slice_size * slice_block_size > attention_slice_rate:
|
||||||
|
slice_2_block_size = split_slice_size * shape_three * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
do_split_2 = True
|
||||||
|
split_2_slice_size = find_slice_size(split_2_slice_size, slice_2_block_size)
|
||||||
|
if split_2_slice_size * slice_2_block_size > attention_slice_rate:
|
||||||
|
slice_3_block_size = split_slice_size * split_2_slice_size * shape_four / 1024 / 1024 * query_element_size
|
||||||
|
do_split_3 = True
|
||||||
|
split_3_slice_size = find_slice_size(split_3_slice_size, slice_3_block_size)
|
||||||
|
|
||||||
|
return do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size
|
||||||
|
|
||||||
|
class SlicedAttnProcessor: # pylint: disable=too-few-public-methods
|
||||||
|
r"""
|
||||||
|
Processor for implementing sliced attention.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
slice_size (`int`, *optional*):
|
||||||
|
The number of steps to compute attention. Uses as many slices as `attention_head_dim // slice_size`, and
|
||||||
|
`attention_head_dim` must be a multiple of the `slice_size`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, slice_size):
|
||||||
|
self.slice_size = slice_size
|
||||||
|
|
||||||
|
def __call__(self, attn: Attention, hidden_states: torch.FloatTensor,
|
||||||
|
encoder_hidden_states=None, attention_mask=None) -> torch.FloatTensor: # pylint: disable=too-many-statements, too-many-locals, too-many-branches
|
||||||
|
|
||||||
|
residual = hidden_states
|
||||||
|
|
||||||
|
input_ndim = hidden_states.ndim
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
batch_size, channel, height, width = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
if attn.group_norm is not None:
|
||||||
|
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
key = attn.to_k(encoder_hidden_states)
|
||||||
|
value = attn.to_v(encoder_hidden_states)
|
||||||
|
key = attn.head_to_batch_dim(key)
|
||||||
|
value = attn.head_to_batch_dim(value)
|
||||||
|
|
||||||
|
batch_size_attention, query_tokens, shape_three = query.shape
|
||||||
|
hidden_states = torch.zeros(
|
||||||
|
(batch_size_attention, query_tokens, dim // attn.heads), device=query.device, dtype=query.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
####################################################################
|
||||||
|
# ARC GPUs can't allocate more than 4GB to a single block, Slice it:
|
||||||
|
_, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_attention_slice_sizes(query.shape, query.element_size(), query.device.type, slice_size=self.slice_size)
|
||||||
|
|
||||||
|
for i in range(batch_size_attention // split_slice_size):
|
||||||
|
start_idx = i * split_slice_size
|
||||||
|
end_idx = (i + 1) * split_slice_size
|
||||||
|
if do_split_2:
|
||||||
|
for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_2 = i2 * split_2_slice_size
|
||||||
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
||||||
|
if do_split_3:
|
||||||
|
for i3 in range(shape_three // split_3_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_3 = i3 * split_3_slice_size
|
||||||
|
end_idx_3 = (i3 + 1) * split_3_slice_size
|
||||||
|
|
||||||
|
query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3]
|
||||||
|
key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3]
|
||||||
|
attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] if attention_mask is not None else None
|
||||||
|
|
||||||
|
attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] = attn_slice
|
||||||
|
del attn_slice
|
||||||
|
else:
|
||||||
|
query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2]
|
||||||
|
key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2]
|
||||||
|
attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2] if attention_mask is not None else None
|
||||||
|
|
||||||
|
attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = attn_slice
|
||||||
|
del attn_slice
|
||||||
|
torch.xpu.synchronize(query.device)
|
||||||
|
else:
|
||||||
|
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)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx] = attn_slice
|
||||||
|
del 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)
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||||
|
|
||||||
|
if attn.residual_connection:
|
||||||
|
hidden_states = hidden_states + residual
|
||||||
|
|
||||||
|
hidden_states = hidden_states / attn.rescale_output_factor
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class AttnProcessor:
|
||||||
|
r"""
|
||||||
|
Default processor for performing attention-related computations.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __call__(self, attn: Attention, hidden_states: torch.FloatTensor,
|
||||||
|
encoder_hidden_states=None, attention_mask=None,
|
||||||
|
temb=None, scale: float = 1.0) -> torch.Tensor: # pylint: disable=too-many-statements, too-many-locals, too-many-branches
|
||||||
|
|
||||||
|
residual = hidden_states
|
||||||
|
|
||||||
|
args = () if USE_PEFT_BACKEND else (scale,)
|
||||||
|
|
||||||
|
if attn.spatial_norm is not None:
|
||||||
|
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||||
|
|
||||||
|
input_ndim = hidden_states.ndim
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
batch_size, channel, height, width = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
if attn.group_norm is not None:
|
||||||
|
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||||
|
|
||||||
|
query = attn.to_q(hidden_states, *args)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
key = attn.to_k(encoder_hidden_states, *args)
|
||||||
|
value = attn.to_v(encoder_hidden_states, *args)
|
||||||
|
|
||||||
|
query = attn.head_to_batch_dim(query)
|
||||||
|
key = attn.head_to_batch_dim(key)
|
||||||
|
value = attn.head_to_batch_dim(value)
|
||||||
|
|
||||||
|
####################################################################
|
||||||
|
# ARC GPUs can't allocate more than 4GB to a single block, Slice it:
|
||||||
|
batch_size_attention, query_tokens, shape_three = query.shape[0], query.shape[1], query.shape[2]
|
||||||
|
hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype)
|
||||||
|
do_split, do_split_2, do_split_3, split_slice_size, split_2_slice_size, split_3_slice_size = find_attention_slice_sizes(query.shape, query.element_size(), query.device.type)
|
||||||
|
|
||||||
|
if do_split:
|
||||||
|
for i in range(batch_size_attention // split_slice_size):
|
||||||
|
start_idx = i * split_slice_size
|
||||||
|
end_idx = (i + 1) * split_slice_size
|
||||||
|
if do_split_2:
|
||||||
|
for i2 in range(query_tokens // split_2_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_2 = i2 * split_2_slice_size
|
||||||
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
||||||
|
if do_split_3:
|
||||||
|
for i3 in range(shape_three // split_3_slice_size): # pylint: disable=invalid-name
|
||||||
|
start_idx_3 = i3 * split_3_slice_size
|
||||||
|
end_idx_3 = (i3 + 1) * split_3_slice_size
|
||||||
|
|
||||||
|
query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3]
|
||||||
|
key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3]
|
||||||
|
attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] if attention_mask is not None else None
|
||||||
|
|
||||||
|
attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2, start_idx_3:end_idx_3] = attn_slice
|
||||||
|
del attn_slice
|
||||||
|
else:
|
||||||
|
query_slice = query[start_idx:end_idx, start_idx_2:end_idx_2]
|
||||||
|
key_slice = key[start_idx:end_idx, start_idx_2:end_idx_2]
|
||||||
|
attn_mask_slice = attention_mask[start_idx:end_idx, start_idx_2:end_idx_2] if attention_mask is not None else None
|
||||||
|
|
||||||
|
attn_slice = attn.get_attention_scores(query_slice, key_slice, attn_mask_slice)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx, start_idx_2:end_idx_2])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx, start_idx_2:end_idx_2] = attn_slice
|
||||||
|
del attn_slice
|
||||||
|
else:
|
||||||
|
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)
|
||||||
|
del query_slice
|
||||||
|
del key_slice
|
||||||
|
del attn_mask_slice
|
||||||
|
attn_slice = torch.bmm(attn_slice, value[start_idx:end_idx])
|
||||||
|
|
||||||
|
hidden_states[start_idx:end_idx] = attn_slice
|
||||||
|
del attn_slice
|
||||||
|
torch.xpu.synchronize(query.device)
|
||||||
|
else:
|
||||||
|
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||||
|
hidden_states = torch.bmm(attention_probs, value)
|
||||||
|
####################################################################
|
||||||
|
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||||
|
|
||||||
|
# linear proj
|
||||||
|
hidden_states = attn.to_out[0](hidden_states, *args)
|
||||||
|
# dropout
|
||||||
|
hidden_states = attn.to_out[1](hidden_states)
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||||
|
|
||||||
|
if attn.residual_connection:
|
||||||
|
hidden_states = hidden_states + residual
|
||||||
|
|
||||||
|
hidden_states = hidden_states / attn.rescale_output_factor
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
def ipex_diffusers():
|
||||||
|
#ARC GPUs can't allocate more than 4GB to a single block:
|
||||||
|
diffusers.models.attention_processor.SlicedAttnProcessor = SlicedAttnProcessor
|
||||||
|
diffusers.models.attention_processor.AttnProcessor = AttnProcessor
|
||||||
@@ -0,0 +1,183 @@
|
|||||||
|
from collections import defaultdict
|
||||||
|
import torch
|
||||||
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||||
|
import intel_extension_for_pytorch._C as core # pylint: disable=import-error, unused-import
|
||||||
|
|
||||||
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
||||||
|
|
||||||
|
device_supports_fp64 = torch.xpu.has_fp64_dtype()
|
||||||
|
OptState = ipex.cpu.autocast._grad_scaler.OptState
|
||||||
|
_MultiDeviceReplicator = ipex.cpu.autocast._grad_scaler._MultiDeviceReplicator
|
||||||
|
_refresh_per_optimizer_state = ipex.cpu.autocast._grad_scaler._refresh_per_optimizer_state
|
||||||
|
|
||||||
|
def _unscale_grads_(self, optimizer, inv_scale, found_inf, allow_fp16): # pylint: disable=unused-argument
|
||||||
|
per_device_inv_scale = _MultiDeviceReplicator(inv_scale)
|
||||||
|
per_device_found_inf = _MultiDeviceReplicator(found_inf)
|
||||||
|
|
||||||
|
# To set up _amp_foreach_non_finite_check_and_unscale_, split grads by device and dtype.
|
||||||
|
# There could be hundreds of grads, so we'd like to iterate through them just once.
|
||||||
|
# However, we don't know their devices or dtypes in advance.
|
||||||
|
|
||||||
|
# https://stackoverflow.com/questions/5029934/defaultdict-of-defaultdict
|
||||||
|
# Google says mypy struggles with defaultdicts type annotations.
|
||||||
|
per_device_and_dtype_grads = defaultdict(lambda: defaultdict(list)) # type: ignore[var-annotated]
|
||||||
|
# sync grad to master weight
|
||||||
|
if hasattr(optimizer, "sync_grad"):
|
||||||
|
optimizer.sync_grad()
|
||||||
|
with torch.no_grad():
|
||||||
|
for group in optimizer.param_groups:
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if (not allow_fp16) and param.grad.dtype == torch.float16:
|
||||||
|
raise ValueError("Attempting to unscale FP16 gradients.")
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
# is_coalesced() == False means the sparse grad has values with duplicate indices.
|
||||||
|
# coalesce() deduplicates indices and adds all values that have the same index.
|
||||||
|
# For scaled fp16 values, there's a good chance coalescing will cause overflow,
|
||||||
|
# so we should check the coalesced _values().
|
||||||
|
if param.grad.dtype is torch.float16:
|
||||||
|
param.grad = param.grad.coalesce()
|
||||||
|
to_unscale = param.grad._values()
|
||||||
|
else:
|
||||||
|
to_unscale = param.grad
|
||||||
|
|
||||||
|
# -: is there a way to split by device and dtype without appending in the inner loop?
|
||||||
|
to_unscale = to_unscale.to("cpu")
|
||||||
|
per_device_and_dtype_grads[to_unscale.device][
|
||||||
|
to_unscale.dtype
|
||||||
|
].append(to_unscale)
|
||||||
|
|
||||||
|
for _, per_dtype_grads in per_device_and_dtype_grads.items():
|
||||||
|
for grads in per_dtype_grads.values():
|
||||||
|
core._amp_foreach_non_finite_check_and_unscale_(
|
||||||
|
grads,
|
||||||
|
per_device_found_inf.get("cpu"),
|
||||||
|
per_device_inv_scale.get("cpu"),
|
||||||
|
)
|
||||||
|
|
||||||
|
return per_device_found_inf._per_device_tensors
|
||||||
|
|
||||||
|
def unscale_(self, optimizer):
|
||||||
|
"""
|
||||||
|
Divides ("unscales") the optimizer's gradient tensors by the scale factor.
|
||||||
|
:meth:`unscale_` is optional, serving cases where you need to
|
||||||
|
:ref:`modify or inspect gradients<working-with-unscaled-gradients>`
|
||||||
|
between the backward pass(es) and :meth:`step`.
|
||||||
|
If :meth:`unscale_` is not called explicitly, gradients will be unscaled automatically during :meth:`step`.
|
||||||
|
Simple example, using :meth:`unscale_` to enable clipping of unscaled gradients::
|
||||||
|
...
|
||||||
|
scaler.scale(loss).backward()
|
||||||
|
scaler.unscale_(optimizer)
|
||||||
|
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||||
|
scaler.step(optimizer)
|
||||||
|
scaler.update()
|
||||||
|
Args:
|
||||||
|
optimizer (torch.optim.Optimizer): Optimizer that owns the gradients to be unscaled.
|
||||||
|
.. warning::
|
||||||
|
:meth:`unscale_` should only be called once per optimizer per :meth:`step` call,
|
||||||
|
and only after all gradients for that optimizer's assigned parameters have been accumulated.
|
||||||
|
Calling :meth:`unscale_` twice for a given optimizer between each :meth:`step` triggers a RuntimeError.
|
||||||
|
.. warning::
|
||||||
|
:meth:`unscale_` may unscale sparse gradients out of place, replacing the ``.grad`` attribute.
|
||||||
|
"""
|
||||||
|
if not self._enabled:
|
||||||
|
return
|
||||||
|
|
||||||
|
self._check_scale_growth_tracker("unscale_")
|
||||||
|
|
||||||
|
optimizer_state = self._per_optimizer_states[id(optimizer)]
|
||||||
|
|
||||||
|
if optimizer_state["stage"] is OptState.UNSCALED: # pylint: disable=no-else-raise
|
||||||
|
raise RuntimeError(
|
||||||
|
"unscale_() has already been called on this optimizer since the last update()."
|
||||||
|
)
|
||||||
|
elif optimizer_state["stage"] is OptState.STEPPED:
|
||||||
|
raise RuntimeError("unscale_() is being called after step().")
|
||||||
|
|
||||||
|
# FP32 division can be imprecise for certain compile options, so we carry out the reciprocal in FP64.
|
||||||
|
assert self._scale is not None
|
||||||
|
if device_supports_fp64:
|
||||||
|
inv_scale = self._scale.double().reciprocal().float()
|
||||||
|
else:
|
||||||
|
inv_scale = self._scale.to("cpu").double().reciprocal().float().to(self._scale.device)
|
||||||
|
found_inf = torch.full(
|
||||||
|
(1,), 0.0, dtype=torch.float32, device=self._scale.device
|
||||||
|
)
|
||||||
|
|
||||||
|
optimizer_state["found_inf_per_device"] = self._unscale_grads_(
|
||||||
|
optimizer, inv_scale, found_inf, False
|
||||||
|
)
|
||||||
|
optimizer_state["stage"] = OptState.UNSCALED
|
||||||
|
|
||||||
|
def update(self, new_scale=None):
|
||||||
|
"""
|
||||||
|
Updates the scale factor.
|
||||||
|
If any optimizer steps were skipped the scale is multiplied by ``backoff_factor``
|
||||||
|
to reduce it. If ``growth_interval`` unskipped iterations occurred consecutively,
|
||||||
|
the scale is multiplied by ``growth_factor`` to increase it.
|
||||||
|
Passing ``new_scale`` sets the new scale value manually. (``new_scale`` is not
|
||||||
|
used directly, it's used to fill GradScaler's internal scale tensor. So if
|
||||||
|
``new_scale`` was a tensor, later in-place changes to that tensor will not further
|
||||||
|
affect the scale GradScaler uses internally.)
|
||||||
|
Args:
|
||||||
|
new_scale (float or :class:`torch.FloatTensor`, optional, default=None): New scale factor.
|
||||||
|
.. warning::
|
||||||
|
:meth:`update` should only be called at the end of the iteration, after ``scaler.step(optimizer)`` has
|
||||||
|
been invoked for all optimizers used this iteration.
|
||||||
|
"""
|
||||||
|
if not self._enabled:
|
||||||
|
return
|
||||||
|
|
||||||
|
_scale, _growth_tracker = self._check_scale_growth_tracker("update")
|
||||||
|
|
||||||
|
if new_scale is not None:
|
||||||
|
# Accept a new user-defined scale.
|
||||||
|
if isinstance(new_scale, float):
|
||||||
|
self._scale.fill_(new_scale) # type: ignore[union-attr]
|
||||||
|
else:
|
||||||
|
reason = "new_scale should be a float or a 1-element torch.FloatTensor with requires_grad=False."
|
||||||
|
assert isinstance(new_scale, torch.FloatTensor), reason # type: ignore[attr-defined]
|
||||||
|
assert new_scale.numel() == 1, reason
|
||||||
|
assert new_scale.requires_grad is False, reason
|
||||||
|
self._scale.copy_(new_scale) # type: ignore[union-attr]
|
||||||
|
else:
|
||||||
|
# Consume shared inf/nan data collected from optimizers to update the scale.
|
||||||
|
# If all found_inf tensors are on the same device as self._scale, this operation is asynchronous.
|
||||||
|
found_infs = [
|
||||||
|
found_inf.to(device="cpu", non_blocking=True)
|
||||||
|
for state in self._per_optimizer_states.values()
|
||||||
|
for found_inf in state["found_inf_per_device"].values()
|
||||||
|
]
|
||||||
|
|
||||||
|
assert len(found_infs) > 0, "No inf checks were recorded prior to update."
|
||||||
|
|
||||||
|
found_inf_combined = found_infs[0]
|
||||||
|
if len(found_infs) > 1:
|
||||||
|
for i in range(1, len(found_infs)):
|
||||||
|
found_inf_combined += found_infs[i]
|
||||||
|
|
||||||
|
to_device = _scale.device
|
||||||
|
_scale = _scale.to("cpu")
|
||||||
|
_growth_tracker = _growth_tracker.to("cpu")
|
||||||
|
|
||||||
|
core._amp_update_scale_(
|
||||||
|
_scale,
|
||||||
|
_growth_tracker,
|
||||||
|
found_inf_combined,
|
||||||
|
self._growth_factor,
|
||||||
|
self._backoff_factor,
|
||||||
|
self._growth_interval,
|
||||||
|
)
|
||||||
|
|
||||||
|
_scale = _scale.to(to_device)
|
||||||
|
_growth_tracker = _growth_tracker.to(to_device)
|
||||||
|
# To prepare for next iteration, clear the data collected from optimizers this iteration.
|
||||||
|
self._per_optimizer_states = defaultdict(_refresh_per_optimizer_state)
|
||||||
|
|
||||||
|
def gradscaler_init():
|
||||||
|
torch.xpu.amp.GradScaler = ipex.cpu.autocast._grad_scaler.GradScaler
|
||||||
|
torch.xpu.amp.GradScaler._unscale_grads_ = _unscale_grads_
|
||||||
|
torch.xpu.amp.GradScaler.unscale_ = unscale_
|
||||||
|
torch.xpu.amp.GradScaler.update = update
|
||||||
|
return torch.xpu.amp.GradScaler
|
||||||
@@ -0,0 +1,313 @@
|
|||||||
|
import os
|
||||||
|
from functools import wraps
|
||||||
|
from contextlib import nullcontext
|
||||||
|
import torch
|
||||||
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
device_supports_fp64 = torch.xpu.has_fp64_dtype()
|
||||||
|
|
||||||
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long, unnecessary-lambda, no-else-return
|
||||||
|
|
||||||
|
class DummyDataParallel(torch.nn.Module): # pylint: disable=missing-class-docstring, unused-argument, too-few-public-methods
|
||||||
|
def __new__(cls, module, device_ids=None, output_device=None, dim=0): # pylint: disable=unused-argument
|
||||||
|
if isinstance(device_ids, list) and len(device_ids) > 1:
|
||||||
|
print("IPEX backend doesn't support DataParallel on multiple XPU devices")
|
||||||
|
return module.to("xpu")
|
||||||
|
|
||||||
|
def return_null_context(*args, **kwargs): # pylint: disable=unused-argument
|
||||||
|
return nullcontext()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_cuda(self):
|
||||||
|
return self.device.type == 'xpu' or self.device.type == 'cuda'
|
||||||
|
|
||||||
|
def check_device(device):
|
||||||
|
return bool((isinstance(device, torch.device) and device.type == "cuda") or (isinstance(device, str) and "cuda" in device) or isinstance(device, int))
|
||||||
|
|
||||||
|
def return_xpu(device):
|
||||||
|
return f"xpu:{device.split(':')[-1]}" if isinstance(device, str) and ":" in device else f"xpu:{device}" if isinstance(device, int) else torch.device("xpu") if isinstance(device, torch.device) else "xpu"
|
||||||
|
|
||||||
|
|
||||||
|
# Autocast
|
||||||
|
original_autocast_init = torch.amp.autocast_mode.autocast.__init__
|
||||||
|
@wraps(torch.amp.autocast_mode.autocast.__init__)
|
||||||
|
def autocast_init(self, device_type, dtype=None, enabled=True, cache_enabled=None):
|
||||||
|
if device_type == "cuda":
|
||||||
|
return original_autocast_init(self, device_type="xpu", dtype=dtype, enabled=enabled, cache_enabled=cache_enabled)
|
||||||
|
else:
|
||||||
|
return original_autocast_init(self, device_type=device_type, dtype=dtype, enabled=enabled, cache_enabled=cache_enabled)
|
||||||
|
|
||||||
|
# Latent Antialias CPU Offload:
|
||||||
|
original_interpolate = torch.nn.functional.interpolate
|
||||||
|
@wraps(torch.nn.functional.interpolate)
|
||||||
|
def interpolate(tensor, size=None, scale_factor=None, mode='nearest', align_corners=None, recompute_scale_factor=None, antialias=False): # pylint: disable=too-many-arguments
|
||||||
|
if antialias or align_corners is not None or mode == 'bicubic':
|
||||||
|
return_device = tensor.device
|
||||||
|
return_dtype = tensor.dtype
|
||||||
|
return original_interpolate(tensor.to("cpu", dtype=torch.float32), size=size, scale_factor=scale_factor, mode=mode,
|
||||||
|
align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias).to(return_device, dtype=return_dtype)
|
||||||
|
else:
|
||||||
|
return original_interpolate(tensor, size=size, scale_factor=scale_factor, mode=mode,
|
||||||
|
align_corners=align_corners, recompute_scale_factor=recompute_scale_factor, antialias=antialias)
|
||||||
|
|
||||||
|
|
||||||
|
# Diffusers Float64 (Alchemist GPUs doesn't support 64 bit):
|
||||||
|
original_from_numpy = torch.from_numpy
|
||||||
|
@wraps(torch.from_numpy)
|
||||||
|
def from_numpy(ndarray):
|
||||||
|
if ndarray.dtype == float:
|
||||||
|
return original_from_numpy(ndarray.astype('float32'))
|
||||||
|
else:
|
||||||
|
return original_from_numpy(ndarray)
|
||||||
|
|
||||||
|
original_as_tensor = torch.as_tensor
|
||||||
|
@wraps(torch.as_tensor)
|
||||||
|
def as_tensor(data, dtype=None, device=None):
|
||||||
|
if check_device(device):
|
||||||
|
device = return_xpu(device)
|
||||||
|
if isinstance(data, np.ndarray) and data.dtype == float and not (
|
||||||
|
(isinstance(device, torch.device) and device.type == "cpu") or (isinstance(device, str) and "cpu" in device)):
|
||||||
|
return original_as_tensor(data, dtype=torch.float32, device=device)
|
||||||
|
else:
|
||||||
|
return original_as_tensor(data, dtype=dtype, device=device)
|
||||||
|
|
||||||
|
|
||||||
|
if device_supports_fp64 and os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None) is None:
|
||||||
|
original_torch_bmm = torch.bmm
|
||||||
|
original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention
|
||||||
|
else:
|
||||||
|
# 32 bit attention workarounds for Alchemist:
|
||||||
|
try:
|
||||||
|
from .attention import torch_bmm_32_bit as original_torch_bmm
|
||||||
|
from .attention import scaled_dot_product_attention_32_bit as original_scaled_dot_product_attention
|
||||||
|
except Exception: # pylint: disable=broad-exception-caught
|
||||||
|
original_torch_bmm = torch.bmm
|
||||||
|
original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention
|
||||||
|
|
||||||
|
|
||||||
|
# Data Type Errors:
|
||||||
|
@wraps(torch.bmm)
|
||||||
|
def torch_bmm(input, mat2, *, out=None):
|
||||||
|
if input.dtype != mat2.dtype:
|
||||||
|
mat2 = mat2.to(input.dtype)
|
||||||
|
return original_torch_bmm(input, mat2, out=out)
|
||||||
|
|
||||||
|
@wraps(torch.nn.functional.scaled_dot_product_attention)
|
||||||
|
def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False):
|
||||||
|
if query.dtype != key.dtype:
|
||||||
|
key = key.to(dtype=query.dtype)
|
||||||
|
if query.dtype != value.dtype:
|
||||||
|
value = value.to(dtype=query.dtype)
|
||||||
|
if attn_mask is not None and query.dtype != attn_mask.dtype:
|
||||||
|
attn_mask = attn_mask.to(dtype=query.dtype)
|
||||||
|
return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal)
|
||||||
|
|
||||||
|
# A1111 FP16
|
||||||
|
original_functional_group_norm = torch.nn.functional.group_norm
|
||||||
|
@wraps(torch.nn.functional.group_norm)
|
||||||
|
def functional_group_norm(input, num_groups, weight=None, bias=None, eps=1e-05):
|
||||||
|
if weight is not None and input.dtype != weight.data.dtype:
|
||||||
|
input = input.to(dtype=weight.data.dtype)
|
||||||
|
if bias is not None and weight is not None and bias.data.dtype != weight.data.dtype:
|
||||||
|
bias.data = bias.data.to(dtype=weight.data.dtype)
|
||||||
|
return original_functional_group_norm(input, num_groups, weight=weight, bias=bias, eps=eps)
|
||||||
|
|
||||||
|
# A1111 BF16
|
||||||
|
original_functional_layer_norm = torch.nn.functional.layer_norm
|
||||||
|
@wraps(torch.nn.functional.layer_norm)
|
||||||
|
def functional_layer_norm(input, normalized_shape, weight=None, bias=None, eps=1e-05):
|
||||||
|
if weight is not None and input.dtype != weight.data.dtype:
|
||||||
|
input = input.to(dtype=weight.data.dtype)
|
||||||
|
if bias is not None and weight is not None and bias.data.dtype != weight.data.dtype:
|
||||||
|
bias.data = bias.data.to(dtype=weight.data.dtype)
|
||||||
|
return original_functional_layer_norm(input, normalized_shape, weight=weight, bias=bias, eps=eps)
|
||||||
|
|
||||||
|
# Training
|
||||||
|
original_functional_linear = torch.nn.functional.linear
|
||||||
|
@wraps(torch.nn.functional.linear)
|
||||||
|
def functional_linear(input, weight, bias=None):
|
||||||
|
if input.dtype != weight.data.dtype:
|
||||||
|
input = input.to(dtype=weight.data.dtype)
|
||||||
|
if bias is not None and bias.data.dtype != weight.data.dtype:
|
||||||
|
bias.data = bias.data.to(dtype=weight.data.dtype)
|
||||||
|
return original_functional_linear(input, weight, bias=bias)
|
||||||
|
|
||||||
|
original_functional_conv2d = torch.nn.functional.conv2d
|
||||||
|
@wraps(torch.nn.functional.conv2d)
|
||||||
|
def functional_conv2d(input, weight, bias=None, stride=1, padding=0, dilation=1, groups=1):
|
||||||
|
if input.dtype != weight.data.dtype:
|
||||||
|
input = input.to(dtype=weight.data.dtype)
|
||||||
|
if bias is not None and bias.data.dtype != weight.data.dtype:
|
||||||
|
bias.data = bias.data.to(dtype=weight.data.dtype)
|
||||||
|
return original_functional_conv2d(input, weight, bias=bias, stride=stride, padding=padding, dilation=dilation, groups=groups)
|
||||||
|
|
||||||
|
# A1111 Embedding BF16
|
||||||
|
original_torch_cat = torch.cat
|
||||||
|
@wraps(torch.cat)
|
||||||
|
def torch_cat(tensor, *args, **kwargs):
|
||||||
|
if len(tensor) == 3 and (tensor[0].dtype != tensor[1].dtype or tensor[2].dtype != tensor[1].dtype):
|
||||||
|
return original_torch_cat([tensor[0].to(tensor[1].dtype), tensor[1], tensor[2].to(tensor[1].dtype)], *args, **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_cat(tensor, *args, **kwargs)
|
||||||
|
|
||||||
|
# SwinIR BF16:
|
||||||
|
original_functional_pad = torch.nn.functional.pad
|
||||||
|
@wraps(torch.nn.functional.pad)
|
||||||
|
def functional_pad(input, pad, mode='constant', value=None):
|
||||||
|
if mode == 'reflect' and input.dtype == torch.bfloat16:
|
||||||
|
return original_functional_pad(input.to(torch.float32), pad, mode=mode, value=value).to(dtype=torch.bfloat16)
|
||||||
|
else:
|
||||||
|
return original_functional_pad(input, pad, mode=mode, value=value)
|
||||||
|
|
||||||
|
|
||||||
|
original_torch_tensor = torch.tensor
|
||||||
|
@wraps(torch.tensor)
|
||||||
|
def torch_tensor(data, *args, dtype=None, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
device = return_xpu(device)
|
||||||
|
if not device_supports_fp64:
|
||||||
|
if (isinstance(device, torch.device) and device.type == "xpu") or (isinstance(device, str) and "xpu" in device):
|
||||||
|
if dtype == torch.float64:
|
||||||
|
dtype = torch.float32
|
||||||
|
elif dtype is None and (hasattr(data, "dtype") and (data.dtype == torch.float64 or data.dtype == float)):
|
||||||
|
dtype = torch.float32
|
||||||
|
return original_torch_tensor(data, *args, dtype=dtype, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_Tensor_to = torch.Tensor.to
|
||||||
|
@wraps(torch.Tensor.to)
|
||||||
|
def Tensor_to(self, device=None, *args, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_Tensor_to(self, return_xpu(device), *args, **kwargs)
|
||||||
|
else:
|
||||||
|
return original_Tensor_to(self, device, *args, **kwargs)
|
||||||
|
|
||||||
|
original_Tensor_cuda = torch.Tensor.cuda
|
||||||
|
@wraps(torch.Tensor.cuda)
|
||||||
|
def Tensor_cuda(self, device=None, *args, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_Tensor_cuda(self, return_xpu(device), *args, **kwargs)
|
||||||
|
else:
|
||||||
|
return original_Tensor_cuda(self, device, *args, **kwargs)
|
||||||
|
|
||||||
|
original_Tensor_pin_memory = torch.Tensor.pin_memory
|
||||||
|
@wraps(torch.Tensor.pin_memory)
|
||||||
|
def Tensor_pin_memory(self, device=None, *args, **kwargs):
|
||||||
|
if device is None:
|
||||||
|
device = "xpu"
|
||||||
|
if check_device(device):
|
||||||
|
return original_Tensor_pin_memory(self, return_xpu(device), *args, **kwargs)
|
||||||
|
else:
|
||||||
|
return original_Tensor_pin_memory(self, device, *args, **kwargs)
|
||||||
|
|
||||||
|
original_UntypedStorage_init = torch.UntypedStorage.__init__
|
||||||
|
@wraps(torch.UntypedStorage.__init__)
|
||||||
|
def UntypedStorage_init(*args, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_UntypedStorage_init(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_UntypedStorage_init(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_UntypedStorage_cuda = torch.UntypedStorage.cuda
|
||||||
|
@wraps(torch.UntypedStorage.cuda)
|
||||||
|
def UntypedStorage_cuda(self, device=None, *args, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_UntypedStorage_cuda(self, return_xpu(device), *args, **kwargs)
|
||||||
|
else:
|
||||||
|
return original_UntypedStorage_cuda(self, device, *args, **kwargs)
|
||||||
|
|
||||||
|
original_torch_empty = torch.empty
|
||||||
|
@wraps(torch.empty)
|
||||||
|
def torch_empty(*args, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_empty(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_empty(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_torch_randn = torch.randn
|
||||||
|
@wraps(torch.randn)
|
||||||
|
def torch_randn(*args, device=None, dtype=None, **kwargs):
|
||||||
|
if dtype == bytes:
|
||||||
|
dtype = None
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_randn(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_randn(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_torch_ones = torch.ones
|
||||||
|
@wraps(torch.ones)
|
||||||
|
def torch_ones(*args, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_ones(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_ones(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_torch_zeros = torch.zeros
|
||||||
|
@wraps(torch.zeros)
|
||||||
|
def torch_zeros(*args, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_zeros(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_zeros(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_torch_linspace = torch.linspace
|
||||||
|
@wraps(torch.linspace)
|
||||||
|
def torch_linspace(*args, device=None, **kwargs):
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_linspace(*args, device=return_xpu(device), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_linspace(*args, device=device, **kwargs)
|
||||||
|
|
||||||
|
original_torch_Generator = torch.Generator
|
||||||
|
@wraps(torch.Generator)
|
||||||
|
def torch_Generator(device=None):
|
||||||
|
if check_device(device):
|
||||||
|
return original_torch_Generator(return_xpu(device))
|
||||||
|
else:
|
||||||
|
return original_torch_Generator(device)
|
||||||
|
|
||||||
|
original_torch_load = torch.load
|
||||||
|
@wraps(torch.load)
|
||||||
|
def torch_load(f, map_location=None, *args, **kwargs):
|
||||||
|
if map_location is None:
|
||||||
|
map_location = "xpu"
|
||||||
|
if check_device(map_location):
|
||||||
|
return original_torch_load(f, *args, map_location=return_xpu(map_location), **kwargs)
|
||||||
|
else:
|
||||||
|
return original_torch_load(f, *args, map_location=map_location, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
# Hijack Functions:
|
||||||
|
def ipex_hijacks():
|
||||||
|
torch.tensor = torch_tensor
|
||||||
|
torch.Tensor.to = Tensor_to
|
||||||
|
torch.Tensor.cuda = Tensor_cuda
|
||||||
|
torch.Tensor.pin_memory = Tensor_pin_memory
|
||||||
|
torch.UntypedStorage.__init__ = UntypedStorage_init
|
||||||
|
torch.UntypedStorage.cuda = UntypedStorage_cuda
|
||||||
|
torch.empty = torch_empty
|
||||||
|
torch.randn = torch_randn
|
||||||
|
torch.ones = torch_ones
|
||||||
|
torch.zeros = torch_zeros
|
||||||
|
torch.linspace = torch_linspace
|
||||||
|
torch.Generator = torch_Generator
|
||||||
|
torch.load = torch_load
|
||||||
|
|
||||||
|
torch.backends.cuda.sdp_kernel = return_null_context
|
||||||
|
torch.nn.DataParallel = DummyDataParallel
|
||||||
|
torch.UntypedStorage.is_cuda = is_cuda
|
||||||
|
torch.amp.autocast_mode.autocast.__init__ = autocast_init
|
||||||
|
|
||||||
|
torch.nn.functional.scaled_dot_product_attention = scaled_dot_product_attention
|
||||||
|
torch.nn.functional.group_norm = functional_group_norm
|
||||||
|
torch.nn.functional.layer_norm = functional_layer_norm
|
||||||
|
torch.nn.functional.linear = functional_linear
|
||||||
|
torch.nn.functional.conv2d = functional_conv2d
|
||||||
|
torch.nn.functional.interpolate = interpolate
|
||||||
|
torch.nn.functional.pad = functional_pad
|
||||||
|
|
||||||
|
torch.bmm = torch_bmm
|
||||||
|
torch.cat = torch_cat
|
||||||
|
if not device_supports_fp64:
|
||||||
|
torch.from_numpy = from_numpy
|
||||||
|
torch.as_tensor = as_tensor
|
||||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,337 @@
|
|||||||
|
# based on https://github.com/Stability-AI/ModelSpec
|
||||||
|
import datetime
|
||||||
|
import hashlib
|
||||||
|
from io import BytesIO
|
||||||
|
import os
|
||||||
|
from typing import List, Optional, Tuple, Union
|
||||||
|
import safetensors
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
r"""
|
||||||
|
# Metadata Example
|
||||||
|
metadata = {
|
||||||
|
# === Must ===
|
||||||
|
"modelspec.sai_model_spec": "1.0.0", # Required version ID for the spec
|
||||||
|
"modelspec.architecture": "stable-diffusion-xl-v1-base", # Architecture, reference the ID of the original model of the arch to match the ID
|
||||||
|
"modelspec.implementation": "sgm",
|
||||||
|
"modelspec.title": "Example Model Version 1.0", # Clean, human-readable title. May use your own phrasing/language/etc
|
||||||
|
# === Should ===
|
||||||
|
"modelspec.author": "Example Corp", # Your name or company name
|
||||||
|
"modelspec.description": "This is my example model to show you how to do it!", # Describe the model in your own words/language/etc. Focus on what users need to know
|
||||||
|
"modelspec.date": "2023-07-20", # ISO-8601 compliant date of when the model was created
|
||||||
|
# === Can ===
|
||||||
|
"modelspec.license": "ExampleLicense-1.0", # eg CreativeML Open RAIL, etc.
|
||||||
|
"modelspec.usage_hint": "Use keyword 'example'" # In your own language, very short hints about how the user should use the model
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
|
||||||
|
BASE_METADATA = {
|
||||||
|
# === Must ===
|
||||||
|
"modelspec.sai_model_spec": "1.0.0", # Required version ID for the spec
|
||||||
|
"modelspec.architecture": None,
|
||||||
|
"modelspec.implementation": None,
|
||||||
|
"modelspec.title": None,
|
||||||
|
"modelspec.resolution": None,
|
||||||
|
# === Should ===
|
||||||
|
"modelspec.description": None,
|
||||||
|
"modelspec.author": None,
|
||||||
|
"modelspec.date": None,
|
||||||
|
# === Can ===
|
||||||
|
"modelspec.license": None,
|
||||||
|
"modelspec.tags": None,
|
||||||
|
"modelspec.merged_from": None,
|
||||||
|
"modelspec.prediction_type": None,
|
||||||
|
"modelspec.timestep_range": None,
|
||||||
|
"modelspec.encoder_layer": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
# 別に使うやつだけ定義
|
||||||
|
MODELSPEC_TITLE = "modelspec.title"
|
||||||
|
|
||||||
|
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_FLUX_1_DEV = "flux-1-dev"
|
||||||
|
ARCH_FLUX_1_UNKNOWN = "flux-1"
|
||||||
|
|
||||||
|
ADAPTER_LORA = "lora"
|
||||||
|
ADAPTER_TEXTUAL_INVERSION = "textual-inversion"
|
||||||
|
|
||||||
|
IMPL_STABILITY_AI = "https://github.com/Stability-AI/generative-models"
|
||||||
|
IMPL_COMFY_UI = "https://github.com/comfyanonymous/ComfyUI"
|
||||||
|
IMPL_DIFFUSERS = "diffusers"
|
||||||
|
IMPL_FLUX = "https://github.com/black-forest-labs/flux"
|
||||||
|
|
||||||
|
PRED_TYPE_EPSILON = "epsilon"
|
||||||
|
PRED_TYPE_V = "v"
|
||||||
|
|
||||||
|
|
||||||
|
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 update_hash_sha256(metadata: dict, state_dict: dict):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
def build_metadata(
|
||||||
|
state_dict: Optional[dict],
|
||||||
|
v2: bool,
|
||||||
|
v_parameterization: bool,
|
||||||
|
sdxl: bool,
|
||||||
|
lora: bool,
|
||||||
|
textual_inversion: bool,
|
||||||
|
timestamp: float,
|
||||||
|
title: Optional[str] = None,
|
||||||
|
reso: Optional[Union[int, Tuple[int, int]]] = None,
|
||||||
|
is_stable_diffusion_ckpt: Optional[bool] = None,
|
||||||
|
author: Optional[str] = None,
|
||||||
|
description: Optional[str] = None,
|
||||||
|
license: Optional[str] = None,
|
||||||
|
tags: Optional[str] = None,
|
||||||
|
merged_from: Optional[str] = None,
|
||||||
|
timesteps: Optional[Tuple[int, int]] = None,
|
||||||
|
clip_skip: Optional[int] = None,
|
||||||
|
sd3: Optional[str] = None,
|
||||||
|
flux: Optional[str] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
sd3: only supports "m", flux: only supports "dev"
|
||||||
|
"""
|
||||||
|
# if state_dict is None, hash is not calculated
|
||||||
|
|
||||||
|
metadata = {}
|
||||||
|
metadata.update(BASE_METADATA)
|
||||||
|
|
||||||
|
# TODO メモリを消費せずかつ正しいハッシュ計算の方法がわかったら実装する
|
||||||
|
# if state_dict is not None:
|
||||||
|
# hash = precalculate_safetensors_hashes(state_dict)
|
||||||
|
# metadata["modelspec.hash_sha256"] = hash
|
||||||
|
|
||||||
|
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
|
||||||
|
elif flux is not None:
|
||||||
|
if flux == "dev":
|
||||||
|
arch = ARCH_FLUX_1_DEV
|
||||||
|
else:
|
||||||
|
arch = ARCH_FLUX_1_UNKNOWN
|
||||||
|
elif v2:
|
||||||
|
if v_parameterization:
|
||||||
|
arch = ARCH_SD_V2_768_V
|
||||||
|
else:
|
||||||
|
arch = ARCH_SD_V2_512
|
||||||
|
else:
|
||||||
|
arch = ARCH_SD_V1
|
||||||
|
|
||||||
|
if lora:
|
||||||
|
arch += f"/{ADAPTER_LORA}"
|
||||||
|
elif textual_inversion:
|
||||||
|
arch += f"/{ADAPTER_TEXTUAL_INVERSION}"
|
||||||
|
|
||||||
|
metadata["modelspec.architecture"] = arch
|
||||||
|
|
||||||
|
if not lora and not textual_inversion and is_stable_diffusion_ckpt is None:
|
||||||
|
is_stable_diffusion_ckpt = True # default is stable diffusion ckpt if not lora and not textual_inversion
|
||||||
|
|
||||||
|
if flux is not None:
|
||||||
|
# Flux
|
||||||
|
impl = IMPL_FLUX
|
||||||
|
elif (lora and sdxl) or textual_inversion or is_stable_diffusion_ckpt:
|
||||||
|
# Stable Diffusion ckpt, TI, SDXL LoRA
|
||||||
|
impl = IMPL_STABILITY_AI
|
||||||
|
else:
|
||||||
|
# v1/v2 LoRA or Diffusers
|
||||||
|
impl = IMPL_DIFFUSERS
|
||||||
|
metadata["modelspec.implementation"] = impl
|
||||||
|
|
||||||
|
if title is None:
|
||||||
|
if lora:
|
||||||
|
title = "LoRA"
|
||||||
|
elif textual_inversion:
|
||||||
|
title = "TextualInversion"
|
||||||
|
else:
|
||||||
|
title = "Checkpoint"
|
||||||
|
title += f"@{timestamp}"
|
||||||
|
metadata[MODELSPEC_TITLE] = title
|
||||||
|
|
||||||
|
if author is not None:
|
||||||
|
metadata["modelspec.author"] = author
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.author"]
|
||||||
|
|
||||||
|
if description is not None:
|
||||||
|
metadata["modelspec.description"] = description
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.description"]
|
||||||
|
|
||||||
|
if merged_from is not None:
|
||||||
|
metadata["modelspec.merged_from"] = merged_from
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.merged_from"]
|
||||||
|
|
||||||
|
if license is not None:
|
||||||
|
metadata["modelspec.license"] = license
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.license"]
|
||||||
|
|
||||||
|
if tags is not None:
|
||||||
|
metadata["modelspec.tags"] = tags
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.tags"]
|
||||||
|
|
||||||
|
# remove microsecond from time
|
||||||
|
int_ts = int(timestamp)
|
||||||
|
|
||||||
|
# time to iso-8601 compliant date
|
||||||
|
date = datetime.datetime.fromtimestamp(int_ts).isoformat()
|
||||||
|
metadata["modelspec.date"] = date
|
||||||
|
|
||||||
|
if reso is not None:
|
||||||
|
# comma separated to tuple
|
||||||
|
if isinstance(reso, str):
|
||||||
|
reso = tuple(map(int, reso.split(",")))
|
||||||
|
if len(reso) == 1:
|
||||||
|
reso = (reso[0], reso[0])
|
||||||
|
else:
|
||||||
|
# resolution is defined in dataset, so use default
|
||||||
|
if sdxl or sd3 is not None or flux is not None:
|
||||||
|
reso = 1024
|
||||||
|
elif v2 and v_parameterization:
|
||||||
|
reso = 768
|
||||||
|
else:
|
||||||
|
reso = 512
|
||||||
|
if isinstance(reso, int):
|
||||||
|
reso = (reso, reso)
|
||||||
|
|
||||||
|
metadata["modelspec.resolution"] = f"{reso[0]}x{reso[1]}"
|
||||||
|
|
||||||
|
if flux is not None:
|
||||||
|
del metadata["modelspec.prediction_type"]
|
||||||
|
elif v_parameterization:
|
||||||
|
metadata["modelspec.prediction_type"] = PRED_TYPE_V
|
||||||
|
else:
|
||||||
|
metadata["modelspec.prediction_type"] = PRED_TYPE_EPSILON
|
||||||
|
|
||||||
|
if timesteps is not None:
|
||||||
|
if isinstance(timesteps, str) or isinstance(timesteps, int):
|
||||||
|
timesteps = (timesteps, timesteps)
|
||||||
|
if len(timesteps) == 1:
|
||||||
|
timesteps = (timesteps[0], timesteps[0])
|
||||||
|
metadata["modelspec.timestep_range"] = f"{timesteps[0]},{timesteps[1]}"
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.timestep_range"]
|
||||||
|
|
||||||
|
if clip_skip is not None:
|
||||||
|
metadata["modelspec.encoder_layer"] = f"{clip_skip}"
|
||||||
|
else:
|
||||||
|
del metadata["modelspec.encoder_layer"]
|
||||||
|
|
||||||
|
# # assert all values are filled
|
||||||
|
# assert all([v is not None for v in metadata.values()]), metadata
|
||||||
|
if not all([v is not None for v in metadata.values()]):
|
||||||
|
logger.error(f"Internal error: some metadata values are None: {metadata}")
|
||||||
|
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
|
# region utils
|
||||||
|
|
||||||
|
|
||||||
|
def get_title(metadata: dict) -> Optional[str]:
|
||||||
|
return metadata.get(MODELSPEC_TITLE, None)
|
||||||
|
|
||||||
|
|
||||||
|
def load_metadata_from_safetensors(model: str) -> dict:
|
||||||
|
if not model.endswith(".safetensors"):
|
||||||
|
return {}
|
||||||
|
|
||||||
|
with safetensors.safe_open(model, framework="pt") as f:
|
||||||
|
metadata = f.metadata()
|
||||||
|
if metadata is None:
|
||||||
|
metadata = {}
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
|
def build_merged_from(models: List[str]) -> str:
|
||||||
|
def get_title(model: str):
|
||||||
|
metadata = load_metadata_from_safetensors(model)
|
||||||
|
title = metadata.get(MODELSPEC_TITLE, None)
|
||||||
|
if title is None:
|
||||||
|
title = os.path.splitext(os.path.basename(model))[0] # use filename
|
||||||
|
return title
|
||||||
|
|
||||||
|
titles = [get_title(model) for model in models]
|
||||||
|
return ", ".join(titles)
|
||||||
|
|
||||||
|
|
||||||
|
# endregion
|
||||||
|
|
||||||
|
|
||||||
|
r"""
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import argparse
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
from library import train_util
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--ckpt", type=str, required=True)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
print(f"Loading {args.ckpt}")
|
||||||
|
state_dict = load_file(args.ckpt)
|
||||||
|
|
||||||
|
print(f"Calculating metadata")
|
||||||
|
metadata = get(state_dict, False, False, False, False, "sgm", False, False, "title", "date", 256, 1000, 0)
|
||||||
|
print(metadata)
|
||||||
|
del state_dict
|
||||||
|
|
||||||
|
# by reference implementation
|
||||||
|
with open(args.ckpt, mode="rb") as file_data:
|
||||||
|
file_hash = hashlib.sha256()
|
||||||
|
head_len = struct.unpack("Q", file_data.read(8)) # int64 header length prefix
|
||||||
|
header = json.loads(file_data.read(head_len[0])) # header itself, json string
|
||||||
|
content = (
|
||||||
|
file_data.read()
|
||||||
|
) # All other content is tightly packed tensors. Copy to RAM for simplicity, but you can avoid this read with a more careful FS-dependent impl.
|
||||||
|
file_hash.update(content)
|
||||||
|
# ===== Update the hash for modelspec =====
|
||||||
|
by_ref = f"0x{file_hash.hexdigest()}"
|
||||||
|
print(by_ref)
|
||||||
|
print("is same?", by_ref == metadata["modelspec.hash_sha256"])
|
||||||
|
|
||||||
|
"""
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,884 @@
|
|||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
import toml
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from typing import Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import save_file
|
||||||
|
from accelerate import Accelerator, PartialState
|
||||||
|
from tqdm import tqdm
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
from . import sd3_models, sd3_utils, strategy_base, train_util
|
||||||
|
from .device_utils import init_ipex, clean_memory_on_device
|
||||||
|
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
# from transformers import CLIPTokenizer
|
||||||
|
# from library import model_util
|
||||||
|
# , sdxl_model_util, train_util, sdxl_original_unet
|
||||||
|
# from library.sdxl_lpw_stable_diffusion import SdxlStableDiffusionLongPromptWeightingPipeline
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
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],
|
||||||
|
sai_metadata: Optional[dict],
|
||||||
|
save_dtype: Optional[torch.dtype] = None,
|
||||||
|
):
|
||||||
|
r"""
|
||||||
|
Save models to checkpoint file. Only supports unified checkpoint format.
|
||||||
|
"""
|
||||||
|
|
||||||
|
state_dict = {}
|
||||||
|
|
||||||
|
def update_sd(prefix, sd):
|
||||||
|
for k, v in sd.items():
|
||||||
|
key = prefix + k
|
||||||
|
if save_dtype is not None:
|
||||||
|
v = v.detach().clone().to("cpu").to(save_dtype)
|
||||||
|
state_dict[key] = v
|
||||||
|
|
||||||
|
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())
|
||||||
|
|
||||||
|
save_file(state_dict, ckpt_path, metadata=sai_metadata)
|
||||||
|
|
||||||
|
|
||||||
|
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],
|
||||||
|
mmdit: sd3_models.MMDiT,
|
||||||
|
vae: sd3_models.SDVAE,
|
||||||
|
):
|
||||||
|
def sd_saver(ckpt_file, epoch_no, global_step):
|
||||||
|
sai_metadata = train_util.get_sai_model_spec(
|
||||||
|
None, args, False, False, False, is_stable_diffusion_ckpt=True, sd3=mmdit.model_type
|
||||||
|
)
|
||||||
|
save_models(ckpt_file, mmdit, vae, clip_l, clip_g, t5xxl, sai_metadata, save_dtype)
|
||||||
|
|
||||||
|
train_util.save_sd_model_on_train_end_common(args, True, True, epoch, global_step, sd_saver, None)
|
||||||
|
|
||||||
|
|
||||||
|
# epochとstepの保存、メタデータにepoch/stepが含まれ引数が同じになるため、統合している
|
||||||
|
# on_epoch_end: Trueならepoch終了時、Falseならstep経過時
|
||||||
|
def save_sd3_model_on_epoch_end_or_stepwise(
|
||||||
|
args: argparse.Namespace,
|
||||||
|
on_epoch_end: bool,
|
||||||
|
accelerator,
|
||||||
|
save_dtype: torch.dtype,
|
||||||
|
epoch: int,
|
||||||
|
num_train_epochs: int,
|
||||||
|
global_step: int,
|
||||||
|
clip_l: sd3_models.SDClipModel,
|
||||||
|
clip_g: sd3_models.SDXLClipG,
|
||||||
|
t5xxl: Optional[sd3_models.T5XXLModel],
|
||||||
|
mmdit: sd3_models.MMDiT,
|
||||||
|
vae: sd3_models.SDVAE,
|
||||||
|
):
|
||||||
|
def sd_saver(ckpt_file, epoch_no, global_step):
|
||||||
|
sai_metadata = train_util.get_sai_model_spec(
|
||||||
|
None, args, False, False, False, is_stable_diffusion_ckpt=True, sd3=mmdit.model_type
|
||||||
|
)
|
||||||
|
save_models(ckpt_file, mmdit, vae, clip_l, clip_g, t5xxl, sai_metadata, save_dtype)
|
||||||
|
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise_common(
|
||||||
|
args,
|
||||||
|
on_epoch_end,
|
||||||
|
accelerator,
|
||||||
|
True,
|
||||||
|
True,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
sd_saver,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
required=False,
|
||||||
|
help="CLIP-L model path. if not specified, use ckpt's state_dict / CLIP-Lモデルのパス。指定しない場合はckptのstate_dictを使用",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--clip_g",
|
||||||
|
type=str,
|
||||||
|
required=False,
|
||||||
|
help="CLIP-G model path. if not specified, use ckpt's state_dict / CLIP-Gモデルのパス。指定しない場合はckptのstate_dictを使用",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--t5xxl",
|
||||||
|
type=str,
|
||||||
|
required=False,
|
||||||
|
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モデルをチェックポイントに保存する"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_t5xxl", action="store_true", help="save T5-XXL model to checkpoint / T5-XXLモデルをチェックポイントに保存する"
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--t5xxl_device",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="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から)を使用",
|
||||||
|
)
|
||||||
|
|
||||||
|
# copy from Diffusers
|
||||||
|
parser.add_argument(
|
||||||
|
"--weighting_scheme",
|
||||||
|
type=str,
|
||||||
|
default="logit_normal",
|
||||||
|
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap"],
|
||||||
|
)
|
||||||
|
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`.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def verify_sdxl_training_args(args: argparse.Namespace, supportTextEncoderCaching: bool = True):
|
||||||
|
assert not args.v2, "v2 cannot be enabled in SDXL training / SDXL学習ではv2を有効にすることはできません"
|
||||||
|
if args.v_parameterization:
|
||||||
|
logger.warning("v_parameterization will be unexpected / SDXL学習ではv_parameterizationは想定外の動作になります")
|
||||||
|
|
||||||
|
if args.clip_skip is not None:
|
||||||
|
logger.warning("clip_skip will be unexpected / SDXL学習ではclip_skipは動作しません")
|
||||||
|
|
||||||
|
# if args.multires_noise_iterations:
|
||||||
|
# logger.info(
|
||||||
|
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET}, but noise_offset is disabled due to multires_noise_iterations / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されていますが、multires_noise_iterationsが有効になっているためnoise_offsetは無効になります"
|
||||||
|
# )
|
||||||
|
# else:
|
||||||
|
# if args.noise_offset is None:
|
||||||
|
# args.noise_offset = DEFAULT_NOISE_OFFSET
|
||||||
|
# elif args.noise_offset != DEFAULT_NOISE_OFFSET:
|
||||||
|
# logger.info(
|
||||||
|
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET} / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されています"
|
||||||
|
# )
|
||||||
|
# logger.info(f"noise_offset is set to {args.noise_offset} / noise_offsetが{args.noise_offset}に設定されました")
|
||||||
|
|
||||||
|
assert (
|
||||||
|
not hasattr(args, "weighted_captions") or not args.weighted_captions
|
||||||
|
), "weighted_captions cannot be enabled in SDXL training currently / SDXL学習では今のところweighted_captionsを有効にすることはできません"
|
||||||
|
|
||||||
|
if supportTextEncoderCaching:
|
||||||
|
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||||
|
args.cache_text_encoder_outputs = True
|
||||||
|
logger.warning(
|
||||||
|
"cache_text_encoder_outputs is enabled because cache_text_encoder_outputs_to_disk is enabled / "
|
||||||
|
+ "cache_text_encoder_outputs_to_diskが有効になっているためcache_text_encoder_outputsが有効になりました"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# temporary copied from sd3_minimal_inferece.py
|
||||||
|
|
||||||
|
|
||||||
|
def get_sigmas(sampling: sd3_utils.ModelSamplingDiscreteFlow, steps):
|
||||||
|
start = sampling.timestep(sampling.sigma_max)
|
||||||
|
end = sampling.timestep(sampling.sigma_min)
|
||||||
|
timesteps = torch.linspace(start, end, steps)
|
||||||
|
sigs = []
|
||||||
|
for x in range(len(timesteps)):
|
||||||
|
ts = timesteps[x]
|
||||||
|
sigs.append(sampling.sigma(ts))
|
||||||
|
sigs += [0.0]
|
||||||
|
return torch.FloatTensor(sigs)
|
||||||
|
|
||||||
|
|
||||||
|
def max_denoise(model_sampling, sigmas):
|
||||||
|
max_sigma = float(model_sampling.sigma_max)
|
||||||
|
sigma = float(sigmas[0])
|
||||||
|
return math.isclose(max_sigma, sigma, rel_tol=1e-05) or sigma > max_sigma
|
||||||
|
|
||||||
|
|
||||||
|
def do_sample(
|
||||||
|
height: int,
|
||||||
|
width: int,
|
||||||
|
seed: int,
|
||||||
|
cond: Tuple[torch.Tensor, torch.Tensor],
|
||||||
|
neg_cond: Tuple[torch.Tensor, torch.Tensor],
|
||||||
|
mmdit: sd3_models.MMDiT,
|
||||||
|
steps: int,
|
||||||
|
guidance_scale: float,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
device: str,
|
||||||
|
):
|
||||||
|
latent = torch.zeros(1, 16, height // 8, width // 8, device=device)
|
||||||
|
latent = latent.to(dtype).to(device)
|
||||||
|
|
||||||
|
# noise = get_noise(seed, latent).to(device)
|
||||||
|
if seed is not None:
|
||||||
|
generator = torch.manual_seed(seed)
|
||||||
|
noise = (
|
||||||
|
torch.randn(latent.size(), dtype=torch.float32, layout=latent.layout, generator=generator, device="cpu")
|
||||||
|
.to(latent.dtype)
|
||||||
|
.to(device)
|
||||||
|
)
|
||||||
|
|
||||||
|
model_sampling = sd3_utils.ModelSamplingDiscreteFlow(shift=3.0) # 3.0 is for SD3
|
||||||
|
|
||||||
|
sigmas = get_sigmas(model_sampling, steps).to(device)
|
||||||
|
|
||||||
|
noise_scaled = model_sampling.noise_scaling(sigmas[0], noise, latent, max_denoise(model_sampling, sigmas))
|
||||||
|
|
||||||
|
c_crossattn = torch.cat([cond[0], neg_cond[0]]).to(device).to(dtype)
|
||||||
|
y = torch.cat([cond[1], neg_cond[1]]).to(device).to(dtype)
|
||||||
|
|
||||||
|
x = noise_scaled.to(device).to(dtype)
|
||||||
|
# print(x.shape)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for i in tqdm(range(len(sigmas) - 1)):
|
||||||
|
sigma_hat = sigmas[i]
|
||||||
|
|
||||||
|
timestep = model_sampling.timestep(sigma_hat).float()
|
||||||
|
timestep = torch.FloatTensor([timestep, timestep]).to(device)
|
||||||
|
|
||||||
|
x_c_nc = torch.cat([x, x], dim=0)
|
||||||
|
# print(x_c_nc.shape, timestep.shape, c_crossattn.shape, y.shape)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
pos_out, neg_out = batched.chunk(2)
|
||||||
|
denoised = neg_out + (pos_out - neg_out) * guidance_scale
|
||||||
|
# print(denoised.shape)
|
||||||
|
|
||||||
|
# d = to_d(x, sigma_hat, denoised)
|
||||||
|
dims_to_append = x.ndim - sigma_hat.ndim
|
||||||
|
sigma_hat_dims = sigma_hat[(...,) + (None,) * dims_to_append]
|
||||||
|
# print(dims_to_append, x.shape, sigma_hat.shape, denoised.shape, sigma_hat_dims.shape)
|
||||||
|
"""Converts a denoiser output to a Karras ODE derivative."""
|
||||||
|
d = (x - denoised) / sigma_hat_dims
|
||||||
|
|
||||||
|
dt = sigmas[i + 1] - sigma_hat
|
||||||
|
|
||||||
|
# Euler method
|
||||||
|
x = x + d * dt
|
||||||
|
x = x.to(dtype)
|
||||||
|
|
||||||
|
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,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
mmdit,
|
||||||
|
vae,
|
||||||
|
text_encoders,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement=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
|
||||||
|
|
||||||
|
# unwrap unet and text_encoder(s)
|
||||||
|
mmdit = accelerator.unwrap_model(mmdit)
|
||||||
|
text_encoders = [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)
|
||||||
|
|
||||||
|
save_dir = args.output_dir + "/sample"
|
||||||
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
|
||||||
|
# save random state to restore later
|
||||||
|
rng_state = torch.get_rng_state()
|
||||||
|
cuda_rng_state = None
|
||||||
|
try:
|
||||||
|
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||||
|
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():
|
||||||
|
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(
|
||||||
|
accelerator,
|
||||||
|
args,
|
||||||
|
mmdit,
|
||||||
|
text_encoders,
|
||||||
|
vae,
|
||||||
|
save_dir,
|
||||||
|
prompt_dict,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement,
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
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]],
|
||||||
|
vae: sd3_models.SDVAE,
|
||||||
|
save_dir,
|
||||||
|
prompt_dict,
|
||||||
|
epoch,
|
||||||
|
steps,
|
||||||
|
sample_prompts_te_outputs,
|
||||||
|
prompt_replacement,
|
||||||
|
):
|
||||||
|
assert isinstance(prompt_dict, dict)
|
||||||
|
negative_prompt = prompt_dict.get("negative_prompt")
|
||||||
|
sample_steps = prompt_dict.get("sample_steps", 30)
|
||||||
|
width = prompt_dict.get("width", 512)
|
||||||
|
height = prompt_dict.get("height", 512)
|
||||||
|
scale = prompt_dict.get("scale", 7.5)
|
||||||
|
seed = prompt_dict.get("seed")
|
||||||
|
# controlnet_image = prompt_dict.get("controlnet_image")
|
||||||
|
prompt: str = prompt_dict.get("prompt", "")
|
||||||
|
# sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler)
|
||||||
|
|
||||||
|
if prompt_replacement is not None:
|
||||||
|
prompt = prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||||
|
if negative_prompt is not None:
|
||||||
|
negative_prompt = negative_prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||||
|
|
||||||
|
if seed is not None:
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
torch.cuda.manual_seed(seed)
|
||||||
|
else:
|
||||||
|
# True random sample image generation
|
||||||
|
torch.seed()
|
||||||
|
torch.cuda.seed()
|
||||||
|
|
||||||
|
if negative_prompt is None:
|
||||||
|
negative_prompt = ""
|
||||||
|
|
||||||
|
height = max(64, height - height % 8) # round to divisible by 8
|
||||||
|
width = max(64, width - width % 8) # round to divisible by 8
|
||||||
|
logger.info(f"prompt: {prompt}")
|
||||||
|
logger.info(f"negative_prompt: {negative_prompt}")
|
||||||
|
logger.info(f"height: {height}")
|
||||||
|
logger.info(f"width: {width}")
|
||||||
|
logger.info(f"sample_steps: {sample_steps}")
|
||||||
|
logger.info(f"scale: {scale}")
|
||||||
|
# logger.info(f"sample_sampler: {sampler_name}")
|
||||||
|
if seed is not None:
|
||||||
|
logger.info(f"seed: {seed}")
|
||||||
|
|
||||||
|
# encode prompts
|
||||||
|
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])
|
||||||
|
|
||||||
|
lg_out, t5_out, pooled = te_outputs
|
||||||
|
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
|
||||||
|
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))
|
||||||
|
|
||||||
|
# latent to image
|
||||||
|
with torch.no_grad():
|
||||||
|
image = vae.decode(latents)
|
||||||
|
image = image.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)
|
||||||
|
|
||||||
|
image = Image.fromarray(decoded_np)
|
||||||
|
# adding accelerator.wait_for_everyone() here should sync up and ensure that sample images are saved in the same order as the original prompt list
|
||||||
|
# but adding 'enum' to the filename should be enough
|
||||||
|
|
||||||
|
ts_str = time.strftime("%Y%m%d%H%M%S", time.localtime())
|
||||||
|
num_suffix = f"e{epoch:06d}" if epoch is not None else f"{steps:06d}"
|
||||||
|
seed_suffix = "" if seed is None else f"_{seed}"
|
||||||
|
i: int = prompt_dict["enum"]
|
||||||
|
img_filename = f"{'' if args.output_name is None else args.output_name + '_'}{num_suffix}_{i:02d}_{ts_str}{seed_suffix}.png"
|
||||||
|
image.save(os.path.join(save_dir, img_filename))
|
||||||
|
|
||||||
|
# wandb有効時のみログを送信
|
||||||
|
try:
|
||||||
|
wandb_tracker = accelerator.get_tracker("wandb")
|
||||||
|
try:
|
||||||
|
import wandb
|
||||||
|
except ImportError: # 事前に一度確認するのでここはエラー出ないはず
|
||||||
|
raise ImportError("No wandb / wandb がインストールされていないようです")
|
||||||
|
|
||||||
|
wandb_tracker.log({f"sample_{i}": wandb.Image(image)})
|
||||||
|
except: # wandb 無効時
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
# region Diffusers
|
||||||
|
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||||
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
|
from diffusers.utils import BaseOutput
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
|
||||||
|
"""
|
||||||
|
Output class for the scheduler's `step` function output.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||||
|
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||||
|
denoising loop.
|
||||||
|
"""
|
||||||
|
|
||||||
|
prev_sample: torch.FloatTensor
|
||||||
|
|
||||||
|
|
||||||
|
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||||
|
"""
|
||||||
|
Euler scheduler.
|
||||||
|
|
||||||
|
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||||
|
methods the library implements for all schedulers such as loading and saving.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
num_train_timesteps (`int`, defaults to 1000):
|
||||||
|
The number of diffusion steps to train the model.
|
||||||
|
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||||
|
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||||
|
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||||
|
shift (`float`, defaults to 1.0):
|
||||||
|
The shift value for the timestep schedule.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_compatibles = []
|
||||||
|
order = 1
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_train_timesteps: int = 1000,
|
||||||
|
shift: float = 1.0,
|
||||||
|
):
|
||||||
|
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||||
|
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||||
|
|
||||||
|
sigmas = timesteps / num_train_timesteps
|
||||||
|
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||||
|
|
||||||
|
self.timesteps = sigmas * num_train_timesteps
|
||||||
|
|
||||||
|
self._step_index = None
|
||||||
|
self._begin_index = None
|
||||||
|
|
||||||
|
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||||
|
self.sigma_min = self.sigmas[-1].item()
|
||||||
|
self.sigma_max = self.sigmas[0].item()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def step_index(self):
|
||||||
|
"""
|
||||||
|
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||||
|
"""
|
||||||
|
return self._step_index
|
||||||
|
|
||||||
|
@property
|
||||||
|
def begin_index(self):
|
||||||
|
"""
|
||||||
|
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||||
|
"""
|
||||||
|
return self._begin_index
|
||||||
|
|
||||||
|
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||||
|
def set_begin_index(self, begin_index: int = 0):
|
||||||
|
"""
|
||||||
|
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
begin_index (`int`):
|
||||||
|
The begin index for the scheduler.
|
||||||
|
"""
|
||||||
|
self._begin_index = begin_index
|
||||||
|
|
||||||
|
def scale_noise(
|
||||||
|
self,
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
timestep: Union[float, torch.FloatTensor],
|
||||||
|
noise: Optional[torch.FloatTensor] = None,
|
||||||
|
) -> torch.FloatTensor:
|
||||||
|
"""
|
||||||
|
Forward process in flow-matching
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sample (`torch.FloatTensor`):
|
||||||
|
The input sample.
|
||||||
|
timestep (`int`, *optional*):
|
||||||
|
The current timestep in the diffusion chain.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`torch.FloatTensor`:
|
||||||
|
A scaled input sample.
|
||||||
|
"""
|
||||||
|
if self.step_index is None:
|
||||||
|
self._init_step_index(timestep)
|
||||||
|
|
||||||
|
sigma = self.sigmas[self.step_index]
|
||||||
|
sample = sigma * noise + (1.0 - sigma) * sample
|
||||||
|
|
||||||
|
return sample
|
||||||
|
|
||||||
|
def _sigma_to_t(self, sigma):
|
||||||
|
return sigma * self.config.num_train_timesteps
|
||||||
|
|
||||||
|
def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.device] = None):
|
||||||
|
"""
|
||||||
|
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
num_inference_steps (`int`):
|
||||||
|
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||||
|
device (`str` or `torch.device`, *optional*):
|
||||||
|
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||||
|
"""
|
||||||
|
self.num_inference_steps = num_inference_steps
|
||||||
|
|
||||||
|
timesteps = np.linspace(self._sigma_to_t(self.sigma_max), self._sigma_to_t(self.sigma_min), num_inference_steps)
|
||||||
|
|
||||||
|
sigmas = timesteps / self.config.num_train_timesteps
|
||||||
|
sigmas = self.config.shift * sigmas / (1 + (self.config.shift - 1) * sigmas)
|
||||||
|
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
timesteps = sigmas * self.config.num_train_timesteps
|
||||||
|
self.timesteps = timesteps.to(device=device)
|
||||||
|
self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||||
|
|
||||||
|
self._step_index = None
|
||||||
|
self._begin_index = None
|
||||||
|
|
||||||
|
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||||
|
if schedule_timesteps is None:
|
||||||
|
schedule_timesteps = self.timesteps
|
||||||
|
|
||||||
|
indices = (schedule_timesteps == timestep).nonzero()
|
||||||
|
|
||||||
|
# The sigma index that is taken for the **very** first `step`
|
||||||
|
# is always the second index (or the last index if there is only 1)
|
||||||
|
# This way we can ensure we don't accidentally skip a sigma in
|
||||||
|
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||||
|
pos = 1 if len(indices) > 1 else 0
|
||||||
|
|
||||||
|
return indices[pos].item()
|
||||||
|
|
||||||
|
def _init_step_index(self, timestep):
|
||||||
|
if self.begin_index is None:
|
||||||
|
if isinstance(timestep, torch.Tensor):
|
||||||
|
timestep = timestep.to(self.timesteps.device)
|
||||||
|
self._step_index = self.index_for_timestep(timestep)
|
||||||
|
else:
|
||||||
|
self._step_index = self._begin_index
|
||||||
|
|
||||||
|
def step(
|
||||||
|
self,
|
||||||
|
model_output: torch.FloatTensor,
|
||||||
|
timestep: Union[float, torch.FloatTensor],
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
s_churn: float = 0.0,
|
||||||
|
s_tmin: float = 0.0,
|
||||||
|
s_tmax: float = float("inf"),
|
||||||
|
s_noise: float = 1.0,
|
||||||
|
generator: Optional[torch.Generator] = None,
|
||||||
|
return_dict: bool = True,
|
||||||
|
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
|
||||||
|
"""
|
||||||
|
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||||
|
process from the learned model outputs (most often the predicted noise).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_output (`torch.FloatTensor`):
|
||||||
|
The direct output from learned diffusion model.
|
||||||
|
timestep (`float`):
|
||||||
|
The current discrete timestep in the diffusion chain.
|
||||||
|
sample (`torch.FloatTensor`):
|
||||||
|
A current instance of a sample created by the diffusion process.
|
||||||
|
s_churn (`float`):
|
||||||
|
s_tmin (`float`):
|
||||||
|
s_tmax (`float`):
|
||||||
|
s_noise (`float`, defaults to 1.0):
|
||||||
|
Scaling factor for noise added to the sample.
|
||||||
|
generator (`torch.Generator`, *optional*):
|
||||||
|
A random number generator.
|
||||||
|
return_dict (`bool`):
|
||||||
|
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||||
|
tuple.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||||
|
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||||
|
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if isinstance(timestep, int) or isinstance(timestep, torch.IntTensor) or isinstance(timestep, torch.LongTensor):
|
||||||
|
raise ValueError(
|
||||||
|
(
|
||||||
|
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||||
|
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||||
|
" one of the `scheduler.timesteps` as a timestep."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.step_index is None:
|
||||||
|
self._init_step_index(timestep)
|
||||||
|
|
||||||
|
# Upcast to avoid precision issues when computing prev_sample
|
||||||
|
sample = sample.to(torch.float32)
|
||||||
|
|
||||||
|
sigma = self.sigmas[self.step_index]
|
||||||
|
|
||||||
|
gamma = min(s_churn / (len(self.sigmas) - 1), 2**0.5 - 1) if s_tmin <= sigma <= s_tmax else 0.0
|
||||||
|
|
||||||
|
noise = randn_tensor(model_output.shape, dtype=model_output.dtype, device=model_output.device, generator=generator)
|
||||||
|
|
||||||
|
eps = noise * s_noise
|
||||||
|
sigma_hat = sigma * (gamma + 1)
|
||||||
|
|
||||||
|
if gamma > 0:
|
||||||
|
sample = sample + eps * (sigma_hat**2 - sigma**2) ** 0.5
|
||||||
|
|
||||||
|
# 1. compute predicted original sample (x_0) from sigma-scaled predicted noise
|
||||||
|
# NOTE: "original_sample" should not be an expected prediction_type but is left in for
|
||||||
|
# backwards compatibility
|
||||||
|
|
||||||
|
# if self.config.prediction_type == "vector_field":
|
||||||
|
|
||||||
|
denoised = sample - model_output * sigma
|
||||||
|
# 2. Convert to an ODE derivative
|
||||||
|
derivative = (sample - denoised) / sigma_hat
|
||||||
|
|
||||||
|
dt = self.sigmas[self.step_index + 1] - sigma_hat
|
||||||
|
|
||||||
|
prev_sample = sample + derivative * dt
|
||||||
|
# Cast sample back to model compatible dtype
|
||||||
|
prev_sample = prev_sample.to(model_output.dtype)
|
||||||
|
|
||||||
|
# upon completion increase step index by one
|
||||||
|
self._step_index += 1
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (prev_sample,)
|
||||||
|
|
||||||
|
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.config.num_train_timesteps
|
||||||
|
|
||||||
|
|
||||||
|
# endregion
|
||||||
@@ -0,0 +1,530 @@
|
|||||||
|
import math
|
||||||
|
from typing import Dict, Optional, Union, List
|
||||||
|
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 .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
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())
|
||||||
|
|
||||||
|
# similar to model.load_state_dict()
|
||||||
|
if not missing_keys and not unexpected_keys:
|
||||||
|
for k in list(state_dict.keys()):
|
||||||
|
set_module_tensor_to_device(model, k, device, value=state_dict.pop(k), dtype=dtype)
|
||||||
|
return "<All keys matched successfully>"
|
||||||
|
|
||||||
|
# error_msgs
|
||||||
|
error_msgs: List[str] = []
|
||||||
|
if missing_keys:
|
||||||
|
error_msgs.insert(0, "Missing key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in missing_keys)))
|
||||||
|
if unexpected_keys:
|
||||||
|
error_msgs.insert(0, "Unexpected key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in unexpected_keys)))
|
||||||
|
|
||||||
|
raise RuntimeError("Error(s) in loading state_dict for {}:\n\t{}".format(model.__class__.__name__, "\n\t".join(error_msgs)))
|
||||||
|
|
||||||
|
def load_safetensors(path: str, dvc: Union[str, torch.device], disable_mmap: bool = False):
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def load_mmdit(state_dict: Dict, attn_mode: str, dtype: Optional[Union[str, torch.dtype]], device: Union[str, torch.device]):
|
||||||
|
mmdit_sd = {}
|
||||||
|
|
||||||
|
mmdit_prefix = "model.diffusion_model."
|
||||||
|
for k in list(state_dict.keys()):
|
||||||
|
if k.startswith(mmdit_prefix):
|
||||||
|
mmdit_sd[k[len(mmdit_prefix) :]] = state_dict.pop(k)
|
||||||
|
|
||||||
|
# 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, mmdit_sd, device, dtype)
|
||||||
|
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]],
|
||||||
|
device: Union[str, torch.device],
|
||||||
|
disable_mmap: bool = False,
|
||||||
|
):
|
||||||
|
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 "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)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def load_clip_g(
|
||||||
|
state_dict: Dict,
|
||||||
|
clip_g_path: Optional[str],
|
||||||
|
attn_mode: str,
|
||||||
|
clip_dtype: Optional[Union[str, torch.dtype]],
|
||||||
|
device: Union[str, torch.device],
|
||||||
|
disable_mmap: bool = False,
|
||||||
|
):
|
||||||
|
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 "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)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
):
|
||||||
|
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 "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)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
):
|
||||||
|
vae_sd = {}
|
||||||
|
if vae_path:
|
||||||
|
logger.info(f"Loading VAE from {vae_path}...")
|
||||||
|
vae_sd = load_safetensors(vae_path, device, disable_mmap)
|
||||||
|
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)
|
||||||
|
|
||||||
|
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 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:
|
||||||
|
"""Helper for sampler scheduling (ie timestep/sigma calculations) for Discrete Flow models"""
|
||||||
|
|
||||||
|
def __init__(self, shift=1.0):
|
||||||
|
self.shift = shift
|
||||||
|
timesteps = 1000
|
||||||
|
self.sigmas = self.sigma(torch.arange(1, timesteps + 1, 1))
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sigma_min(self):
|
||||||
|
return self.sigmas[0]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sigma_max(self):
|
||||||
|
return self.sigmas[-1]
|
||||||
|
|
||||||
|
def timestep(self, sigma):
|
||||||
|
return sigma * 1000
|
||||||
|
|
||||||
|
def sigma(self, timestep: torch.Tensor):
|
||||||
|
timestep = timestep / 1000.0
|
||||||
|
if self.shift == 1.0:
|
||||||
|
return timestep
|
||||||
|
return self.shift * timestep / (1 + (self.shift - 1) * timestep)
|
||||||
|
|
||||||
|
def calculate_denoised(self, sigma, model_output, model_input):
|
||||||
|
sigma = sigma.view(sigma.shape[:1] + (1,) * (model_output.ndim - 1))
|
||||||
|
return model_input - model_output * sigma
|
||||||
|
|
||||||
|
def noise_scaling(self, sigma, noise, latent_image, max_denoise=False):
|
||||||
|
# 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
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,583 @@
|
|||||||
|
import torch
|
||||||
|
import safetensors
|
||||||
|
from accelerate import init_empty_weights
|
||||||
|
from accelerate.utils.modeling import set_module_tensor_to_device
|
||||||
|
from safetensors.torch import load_file, save_file
|
||||||
|
from transformers import CLIPTextModel, CLIPTextConfig, CLIPTextModelWithProjection, CLIPTokenizer
|
||||||
|
from typing import List
|
||||||
|
from diffusers import AutoencoderKL, EulerDiscreteScheduler, UNet2DConditionModel
|
||||||
|
from library import model_util
|
||||||
|
from library import sdxl_original_unet
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
VAE_SCALE_FACTOR = 0.13025
|
||||||
|
MODEL_VERSION_SDXL_BASE_V1_0 = "sdxl_base_v1-0"
|
||||||
|
|
||||||
|
# Diffusersの設定を読み込むための参照モデル
|
||||||
|
DIFFUSERS_REF_MODEL_ID_SDXL = "stabilityai/stable-diffusion-xl-base-1.0"
|
||||||
|
|
||||||
|
DIFFUSERS_SDXL_UNET_CONFIG = {
|
||||||
|
"act_fn": "silu",
|
||||||
|
"addition_embed_type": "text_time",
|
||||||
|
"addition_embed_type_num_heads": 64,
|
||||||
|
"addition_time_embed_dim": 256,
|
||||||
|
"attention_head_dim": [5, 10, 20],
|
||||||
|
"block_out_channels": [320, 640, 1280],
|
||||||
|
"center_input_sample": False,
|
||||||
|
"class_embed_type": None,
|
||||||
|
"class_embeddings_concat": False,
|
||||||
|
"conv_in_kernel": 3,
|
||||||
|
"conv_out_kernel": 3,
|
||||||
|
"cross_attention_dim": 2048,
|
||||||
|
"cross_attention_norm": None,
|
||||||
|
"down_block_types": ["DownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D"],
|
||||||
|
"downsample_padding": 1,
|
||||||
|
"dual_cross_attention": False,
|
||||||
|
"encoder_hid_dim": None,
|
||||||
|
"encoder_hid_dim_type": None,
|
||||||
|
"flip_sin_to_cos": True,
|
||||||
|
"freq_shift": 0,
|
||||||
|
"in_channels": 4,
|
||||||
|
"layers_per_block": 2,
|
||||||
|
"mid_block_only_cross_attention": None,
|
||||||
|
"mid_block_scale_factor": 1,
|
||||||
|
"mid_block_type": "UNetMidBlock2DCrossAttn",
|
||||||
|
"norm_eps": 1e-05,
|
||||||
|
"norm_num_groups": 32,
|
||||||
|
"num_attention_heads": None,
|
||||||
|
"num_class_embeds": None,
|
||||||
|
"only_cross_attention": False,
|
||||||
|
"out_channels": 4,
|
||||||
|
"projection_class_embeddings_input_dim": 2816,
|
||||||
|
"resnet_out_scale_factor": 1.0,
|
||||||
|
"resnet_skip_time_act": False,
|
||||||
|
"resnet_time_scale_shift": "default",
|
||||||
|
"sample_size": 128,
|
||||||
|
"time_cond_proj_dim": None,
|
||||||
|
"time_embedding_act_fn": None,
|
||||||
|
"time_embedding_dim": None,
|
||||||
|
"time_embedding_type": "positional",
|
||||||
|
"timestep_post_act": None,
|
||||||
|
"transformer_layers_per_block": [1, 2, 10],
|
||||||
|
"up_block_types": ["CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"],
|
||||||
|
"upcast_attention": False,
|
||||||
|
"use_linear_projection": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def convert_sdxl_text_encoder_2_checkpoint(checkpoint, max_length):
|
||||||
|
SDXL_KEY_PREFIX = "conditioner.embedders.1.model."
|
||||||
|
|
||||||
|
# SD2のと、基本的には同じ。logit_scaleを後で使うので、それを追加で返す
|
||||||
|
# logit_scaleはcheckpointの保存時に使用する
|
||||||
|
def convert_key(key):
|
||||||
|
# common conversion
|
||||||
|
key = key.replace(SDXL_KEY_PREFIX + "transformer.", "text_model.encoder.")
|
||||||
|
key = key.replace(SDXL_KEY_PREFIX, "text_model.")
|
||||||
|
|
||||||
|
if "resblocks" in key:
|
||||||
|
# resblocks conversion
|
||||||
|
key = key.replace(".resblocks.", ".layers.")
|
||||||
|
if ".ln_" in key:
|
||||||
|
key = key.replace(".ln_", ".layer_norm")
|
||||||
|
elif ".mlp." in key:
|
||||||
|
key = key.replace(".c_fc.", ".fc1.")
|
||||||
|
key = key.replace(".c_proj.", ".fc2.")
|
||||||
|
elif ".attn.out_proj" in key:
|
||||||
|
key = key.replace(".attn.out_proj.", ".self_attn.out_proj.")
|
||||||
|
elif ".attn.in_proj" in key:
|
||||||
|
key = None # 特殊なので後で処理する
|
||||||
|
else:
|
||||||
|
raise ValueError(f"unexpected key in SD: {key}")
|
||||||
|
elif ".positional_embedding" in key:
|
||||||
|
key = key.replace(".positional_embedding", ".embeddings.position_embedding.weight")
|
||||||
|
elif ".text_projection" in key:
|
||||||
|
key = key.replace("text_model.text_projection", "text_projection.weight")
|
||||||
|
elif ".logit_scale" in key:
|
||||||
|
key = None # 後で処理する
|
||||||
|
elif ".token_embedding" in key:
|
||||||
|
key = key.replace(".token_embedding.weight", ".embeddings.token_embedding.weight")
|
||||||
|
elif ".ln_final" in key:
|
||||||
|
key = key.replace(".ln_final", ".final_layer_norm")
|
||||||
|
# ckpt from comfy has this key: text_model.encoder.text_model.embeddings.position_ids
|
||||||
|
elif ".embeddings.position_ids" in key:
|
||||||
|
key = None # remove this key: position_ids is not used in newer transformers
|
||||||
|
return key
|
||||||
|
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
new_sd = {}
|
||||||
|
for key in keys:
|
||||||
|
new_key = convert_key(key)
|
||||||
|
if new_key is None:
|
||||||
|
continue
|
||||||
|
new_sd[new_key] = checkpoint[key]
|
||||||
|
|
||||||
|
# attnの変換
|
||||||
|
for key in keys:
|
||||||
|
if ".resblocks" in key and ".attn.in_proj_" in key:
|
||||||
|
# 三つに分割
|
||||||
|
values = torch.chunk(checkpoint[key], 3)
|
||||||
|
|
||||||
|
key_suffix = ".weight" if "weight" in key else ".bias"
|
||||||
|
key_pfx = key.replace(SDXL_KEY_PREFIX + "transformer.resblocks.", "text_model.encoder.layers.")
|
||||||
|
key_pfx = key_pfx.replace("_weight", "")
|
||||||
|
key_pfx = key_pfx.replace("_bias", "")
|
||||||
|
key_pfx = key_pfx.replace(".attn.in_proj", ".self_attn.")
|
||||||
|
new_sd[key_pfx + "q_proj" + key_suffix] = values[0]
|
||||||
|
new_sd[key_pfx + "k_proj" + key_suffix] = values[1]
|
||||||
|
new_sd[key_pfx + "v_proj" + key_suffix] = values[2]
|
||||||
|
|
||||||
|
# logit_scale はDiffusersには含まれないが、保存時に戻したいので別途返す
|
||||||
|
logit_scale = checkpoint.get(SDXL_KEY_PREFIX + "logit_scale", None)
|
||||||
|
|
||||||
|
# temporary workaround for text_projection.weight.weight for Playground-v2
|
||||||
|
if "text_projection.weight.weight" in new_sd:
|
||||||
|
logger.info("convert_sdxl_text_encoder_2_checkpoint: convert text_projection.weight.weight to text_projection.weight")
|
||||||
|
new_sd["text_projection.weight"] = new_sd["text_projection.weight.weight"]
|
||||||
|
del new_sd["text_projection.weight.weight"]
|
||||||
|
|
||||||
|
return new_sd, logit_scale
|
||||||
|
|
||||||
|
|
||||||
|
# load state_dict without allocating new tensors
|
||||||
|
def _load_state_dict_on_device(model, state_dict, device, dtype=None):
|
||||||
|
# dtype will use fp32 as default
|
||||||
|
missing_keys = list(model.state_dict().keys() - state_dict.keys())
|
||||||
|
unexpected_keys = list(state_dict.keys() - model.state_dict().keys())
|
||||||
|
|
||||||
|
# similar to model.load_state_dict()
|
||||||
|
if not missing_keys and not unexpected_keys:
|
||||||
|
for k in list(state_dict.keys()):
|
||||||
|
set_module_tensor_to_device(model, k, device, value=state_dict.pop(k), dtype=dtype)
|
||||||
|
return "<All keys matched successfully>"
|
||||||
|
|
||||||
|
# error_msgs
|
||||||
|
error_msgs: List[str] = []
|
||||||
|
if missing_keys:
|
||||||
|
error_msgs.insert(0, "Missing key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in missing_keys)))
|
||||||
|
if unexpected_keys:
|
||||||
|
error_msgs.insert(0, "Unexpected key(s) in state_dict: {}. ".format(", ".join('"{}"'.format(k) for k in unexpected_keys)))
|
||||||
|
|
||||||
|
raise RuntimeError("Error(s) in loading state_dict for {}:\n\t{}".format(model.__class__.__name__, "\n\t".join(error_msgs)))
|
||||||
|
|
||||||
|
|
||||||
|
def load_models_from_sdxl_checkpoint(model_version, ckpt_path, map_location, dtype=None, disable_mmap=False):
|
||||||
|
# model_version is reserved for future use
|
||||||
|
# dtype is used for full_fp16/bf16 integration. Text Encoder will remain fp32, because it runs on CPU when caching
|
||||||
|
|
||||||
|
# Load the state dict
|
||||||
|
if model_util.is_safetensors(ckpt_path):
|
||||||
|
checkpoint = None
|
||||||
|
if disable_mmap:
|
||||||
|
state_dict = safetensors.torch.load(open(ckpt_path, "rb").read())
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
state_dict = load_file(ckpt_path, device=map_location)
|
||||||
|
except:
|
||||||
|
state_dict = load_file(ckpt_path) # prevent device invalid Error
|
||||||
|
epoch = None
|
||||||
|
global_step = None
|
||||||
|
else:
|
||||||
|
checkpoint = torch.load(ckpt_path, map_location=map_location)
|
||||||
|
if "state_dict" in checkpoint:
|
||||||
|
state_dict = checkpoint["state_dict"]
|
||||||
|
epoch = checkpoint.get("epoch", 0)
|
||||||
|
global_step = checkpoint.get("global_step", 0)
|
||||||
|
else:
|
||||||
|
state_dict = checkpoint
|
||||||
|
epoch = 0
|
||||||
|
global_step = 0
|
||||||
|
checkpoint = None
|
||||||
|
|
||||||
|
# U-Net
|
||||||
|
logger.info("building U-Net")
|
||||||
|
with init_empty_weights():
|
||||||
|
unet = sdxl_original_unet.SdxlUNet2DConditionModel()
|
||||||
|
|
||||||
|
logger.info("loading U-Net from checkpoint")
|
||||||
|
unet_sd = {}
|
||||||
|
for k in list(state_dict.keys()):
|
||||||
|
if k.startswith("model.diffusion_model."):
|
||||||
|
unet_sd[k.replace("model.diffusion_model.", "")] = state_dict.pop(k)
|
||||||
|
info = _load_state_dict_on_device(unet, unet_sd, device=map_location, dtype=dtype)
|
||||||
|
logger.info(f"U-Net: {info}")
|
||||||
|
|
||||||
|
# Text Encoders
|
||||||
|
logger.info("building text encoders")
|
||||||
|
|
||||||
|
# Text Encoder 1 is same to Stability AI's SDXL
|
||||||
|
text_model1_cfg = CLIPTextConfig(
|
||||||
|
vocab_size=49408,
|
||||||
|
hidden_size=768,
|
||||||
|
intermediate_size=3072,
|
||||||
|
num_hidden_layers=12,
|
||||||
|
num_attention_heads=12,
|
||||||
|
max_position_embeddings=77,
|
||||||
|
hidden_act="quick_gelu",
|
||||||
|
layer_norm_eps=1e-05,
|
||||||
|
dropout=0.0,
|
||||||
|
attention_dropout=0.0,
|
||||||
|
initializer_range=0.02,
|
||||||
|
initializer_factor=1.0,
|
||||||
|
pad_token_id=1,
|
||||||
|
bos_token_id=0,
|
||||||
|
eos_token_id=2,
|
||||||
|
model_type="clip_text_model",
|
||||||
|
projection_dim=768,
|
||||||
|
# torch_dtype="float32",
|
||||||
|
# transformers_version="4.25.0.dev0",
|
||||||
|
)
|
||||||
|
with init_empty_weights():
|
||||||
|
text_model1 = CLIPTextModel._from_config(text_model1_cfg)
|
||||||
|
|
||||||
|
# Text Encoder 2 is different from Stability AI's SDXL. SDXL uses open clip, but we use the model from HuggingFace.
|
||||||
|
# Note: Tokenizer from HuggingFace is different from SDXL. We must use open clip's tokenizer.
|
||||||
|
text_model2_cfg = CLIPTextConfig(
|
||||||
|
vocab_size=49408,
|
||||||
|
hidden_size=1280,
|
||||||
|
intermediate_size=5120,
|
||||||
|
num_hidden_layers=32,
|
||||||
|
num_attention_heads=20,
|
||||||
|
max_position_embeddings=77,
|
||||||
|
hidden_act="gelu",
|
||||||
|
layer_norm_eps=1e-05,
|
||||||
|
dropout=0.0,
|
||||||
|
attention_dropout=0.0,
|
||||||
|
initializer_range=0.02,
|
||||||
|
initializer_factor=1.0,
|
||||||
|
pad_token_id=1,
|
||||||
|
bos_token_id=0,
|
||||||
|
eos_token_id=2,
|
||||||
|
model_type="clip_text_model",
|
||||||
|
projection_dim=1280,
|
||||||
|
# torch_dtype="float32",
|
||||||
|
# transformers_version="4.25.0.dev0",
|
||||||
|
)
|
||||||
|
with init_empty_weights():
|
||||||
|
text_model2 = CLIPTextModelWithProjection(text_model2_cfg)
|
||||||
|
|
||||||
|
logger.info("loading text encoders from checkpoint")
|
||||||
|
te1_sd = {}
|
||||||
|
te2_sd = {}
|
||||||
|
for k in list(state_dict.keys()):
|
||||||
|
if k.startswith("conditioner.embedders.0.transformer."):
|
||||||
|
te1_sd[k.replace("conditioner.embedders.0.transformer.", "")] = state_dict.pop(k)
|
||||||
|
elif k.startswith("conditioner.embedders.1.model."):
|
||||||
|
te2_sd[k] = state_dict.pop(k)
|
||||||
|
|
||||||
|
# 最新の transformers では position_ids を含むとエラーになるので削除 / remove position_ids for latest transformers
|
||||||
|
if "text_model.embeddings.position_ids" in te1_sd:
|
||||||
|
te1_sd.pop("text_model.embeddings.position_ids")
|
||||||
|
|
||||||
|
info1 = _load_state_dict_on_device(text_model1, te1_sd, device=map_location) # remain fp32
|
||||||
|
logger.info(f"text encoder 1: {info1}")
|
||||||
|
|
||||||
|
converted_sd, logit_scale = convert_sdxl_text_encoder_2_checkpoint(te2_sd, max_length=77)
|
||||||
|
info2 = _load_state_dict_on_device(text_model2, converted_sd, device=map_location) # remain fp32
|
||||||
|
logger.info(f"text encoder 2: {info2}")
|
||||||
|
|
||||||
|
# prepare vae
|
||||||
|
logger.info("building VAE")
|
||||||
|
vae_config = model_util.create_vae_diffusers_config()
|
||||||
|
with init_empty_weights():
|
||||||
|
vae = AutoencoderKL(**vae_config)
|
||||||
|
|
||||||
|
logger.info("loading VAE from checkpoint")
|
||||||
|
converted_vae_checkpoint = model_util.convert_ldm_vae_checkpoint(state_dict, vae_config)
|
||||||
|
info = _load_state_dict_on_device(vae, converted_vae_checkpoint, device=map_location, dtype=dtype)
|
||||||
|
logger.info(f"VAE: {info}")
|
||||||
|
|
||||||
|
ckpt_info = (epoch, global_step) if epoch is not None else None
|
||||||
|
return text_model1, text_model2, vae, unet, logit_scale, ckpt_info
|
||||||
|
|
||||||
|
|
||||||
|
def make_unet_conversion_map():
|
||||||
|
unet_conversion_map_layer = []
|
||||||
|
|
||||||
|
for i in range(3): # num_blocks is 3 in sdxl
|
||||||
|
# loop over downblocks/upblocks
|
||||||
|
for j in range(2):
|
||||||
|
# loop over resnets/attentions for downblocks
|
||||||
|
hf_down_res_prefix = f"down_blocks.{i}.resnets.{j}."
|
||||||
|
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no attention layers in down_blocks.3
|
||||||
|
hf_down_atn_prefix = f"down_blocks.{i}.attentions.{j}."
|
||||||
|
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(3):
|
||||||
|
# loop over resnets/attentions for upblocks
|
||||||
|
hf_up_res_prefix = f"up_blocks.{i}.resnets.{j}."
|
||||||
|
sd_up_res_prefix = f"output_blocks.{3*i + j}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
|
||||||
|
|
||||||
|
# if i > 0: commentout for sdxl
|
||||||
|
# no attention layers in up_blocks.0
|
||||||
|
hf_up_atn_prefix = f"up_blocks.{i}.attentions.{j}."
|
||||||
|
sd_up_atn_prefix = f"output_blocks.{3*i + j}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no downsample in down_blocks.3
|
||||||
|
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0.conv."
|
||||||
|
sd_downsample_prefix = f"input_blocks.{3*(i+1)}.0.op."
|
||||||
|
unet_conversion_map_layer.append((sd_downsample_prefix, hf_downsample_prefix))
|
||||||
|
|
||||||
|
# no upsample in up_blocks.3
|
||||||
|
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
|
||||||
|
sd_upsample_prefix = f"output_blocks.{3*i + 2}.{2}." # change for sdxl
|
||||||
|
unet_conversion_map_layer.append((sd_upsample_prefix, hf_upsample_prefix))
|
||||||
|
|
||||||
|
hf_mid_atn_prefix = "mid_block.attentions.0."
|
||||||
|
sd_mid_atn_prefix = "middle_block.1."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.resnets.{j}."
|
||||||
|
sd_mid_res_prefix = f"middle_block.{2*j}."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
unet_conversion_map_resnet = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("in_layers.0.", "norm1."),
|
||||||
|
("in_layers.2.", "conv1."),
|
||||||
|
("out_layers.0.", "norm2."),
|
||||||
|
("out_layers.3.", "conv2."),
|
||||||
|
("emb_layers.1.", "time_emb_proj."),
|
||||||
|
("skip_connection.", "conv_shortcut."),
|
||||||
|
]
|
||||||
|
|
||||||
|
unet_conversion_map = []
|
||||||
|
for sd, hf in unet_conversion_map_layer:
|
||||||
|
if "resnets" in hf:
|
||||||
|
for sd_res, hf_res in unet_conversion_map_resnet:
|
||||||
|
unet_conversion_map.append((sd + sd_res, hf + hf_res))
|
||||||
|
else:
|
||||||
|
unet_conversion_map.append((sd, hf))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_time_embed_prefix = f"time_embedding.linear_{j+1}."
|
||||||
|
sd_time_embed_prefix = f"time_embed.{j*2}."
|
||||||
|
unet_conversion_map.append((sd_time_embed_prefix, hf_time_embed_prefix))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_label_embed_prefix = f"add_embedding.linear_{j+1}."
|
||||||
|
sd_label_embed_prefix = f"label_emb.0.{j*2}."
|
||||||
|
unet_conversion_map.append((sd_label_embed_prefix, hf_label_embed_prefix))
|
||||||
|
|
||||||
|
unet_conversion_map.append(("input_blocks.0.0.", "conv_in."))
|
||||||
|
unet_conversion_map.append(("out.0.", "conv_norm_out."))
|
||||||
|
unet_conversion_map.append(("out.2.", "conv_out."))
|
||||||
|
|
||||||
|
return unet_conversion_map
|
||||||
|
|
||||||
|
|
||||||
|
def convert_diffusers_unet_state_dict_to_sdxl(du_sd):
|
||||||
|
unet_conversion_map = make_unet_conversion_map()
|
||||||
|
|
||||||
|
conversion_map = {hf: sd for sd, hf in unet_conversion_map}
|
||||||
|
return convert_unet_state_dict(du_sd, conversion_map)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_unet_state_dict(src_sd, conversion_map):
|
||||||
|
converted_sd = {}
|
||||||
|
for src_key, value in src_sd.items():
|
||||||
|
# さすがに全部回すのは時間がかかるので右から要素を削りつつprefixを探す
|
||||||
|
src_key_fragments = src_key.split(".")[:-1] # remove weight/bias
|
||||||
|
while len(src_key_fragments) > 0:
|
||||||
|
src_key_prefix = ".".join(src_key_fragments) + "."
|
||||||
|
if src_key_prefix in conversion_map:
|
||||||
|
converted_prefix = conversion_map[src_key_prefix]
|
||||||
|
converted_key = converted_prefix + src_key[len(src_key_prefix) :]
|
||||||
|
converted_sd[converted_key] = value
|
||||||
|
break
|
||||||
|
src_key_fragments.pop(-1)
|
||||||
|
assert len(src_key_fragments) > 0, f"key {src_key} not found in conversion map"
|
||||||
|
|
||||||
|
return converted_sd
|
||||||
|
|
||||||
|
|
||||||
|
def convert_sdxl_unet_state_dict_to_diffusers(sd):
|
||||||
|
unet_conversion_map = make_unet_conversion_map()
|
||||||
|
|
||||||
|
conversion_dict = {sd: hf for sd, hf in unet_conversion_map}
|
||||||
|
return convert_unet_state_dict(sd, conversion_dict)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_encoder_2_state_dict_to_sdxl(checkpoint, logit_scale):
|
||||||
|
def convert_key(key):
|
||||||
|
# position_idsの除去
|
||||||
|
if ".position_ids" in key:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# common
|
||||||
|
key = key.replace("text_model.encoder.", "transformer.")
|
||||||
|
key = key.replace("text_model.", "")
|
||||||
|
if "layers" in key:
|
||||||
|
# resblocks conversion
|
||||||
|
key = key.replace(".layers.", ".resblocks.")
|
||||||
|
if ".layer_norm" in key:
|
||||||
|
key = key.replace(".layer_norm", ".ln_")
|
||||||
|
elif ".mlp." in key:
|
||||||
|
key = key.replace(".fc1.", ".c_fc.")
|
||||||
|
key = key.replace(".fc2.", ".c_proj.")
|
||||||
|
elif ".self_attn.out_proj" in key:
|
||||||
|
key = key.replace(".self_attn.out_proj.", ".attn.out_proj.")
|
||||||
|
elif ".self_attn." in key:
|
||||||
|
key = None # 特殊なので後で処理する
|
||||||
|
else:
|
||||||
|
raise ValueError(f"unexpected key in DiffUsers model: {key}")
|
||||||
|
elif ".position_embedding" in key:
|
||||||
|
key = key.replace("embeddings.position_embedding.weight", "positional_embedding")
|
||||||
|
elif ".token_embedding" in key:
|
||||||
|
key = key.replace("embeddings.token_embedding.weight", "token_embedding.weight")
|
||||||
|
elif "text_projection" in key: # no dot in key
|
||||||
|
key = key.replace("text_projection.weight", "text_projection")
|
||||||
|
elif "final_layer_norm" in key:
|
||||||
|
key = key.replace("final_layer_norm", "ln_final")
|
||||||
|
return key
|
||||||
|
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
new_sd = {}
|
||||||
|
for key in keys:
|
||||||
|
new_key = convert_key(key)
|
||||||
|
if new_key is None:
|
||||||
|
continue
|
||||||
|
new_sd[new_key] = checkpoint[key]
|
||||||
|
|
||||||
|
# attnの変換
|
||||||
|
for key in keys:
|
||||||
|
if "layers" in key and "q_proj" in key:
|
||||||
|
# 三つを結合
|
||||||
|
key_q = key
|
||||||
|
key_k = key.replace("q_proj", "k_proj")
|
||||||
|
key_v = key.replace("q_proj", "v_proj")
|
||||||
|
|
||||||
|
value_q = checkpoint[key_q]
|
||||||
|
value_k = checkpoint[key_k]
|
||||||
|
value_v = checkpoint[key_v]
|
||||||
|
value = torch.cat([value_q, value_k, value_v])
|
||||||
|
|
||||||
|
new_key = key.replace("text_model.encoder.layers.", "transformer.resblocks.")
|
||||||
|
new_key = new_key.replace(".self_attn.q_proj.", ".attn.in_proj_")
|
||||||
|
new_sd[new_key] = value
|
||||||
|
|
||||||
|
if logit_scale is not None:
|
||||||
|
new_sd["logit_scale"] = logit_scale
|
||||||
|
|
||||||
|
return new_sd
|
||||||
|
|
||||||
|
|
||||||
|
def save_stable_diffusion_checkpoint(
|
||||||
|
output_file,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
epochs,
|
||||||
|
steps,
|
||||||
|
ckpt_info,
|
||||||
|
vae,
|
||||||
|
logit_scale,
|
||||||
|
metadata,
|
||||||
|
save_dtype=None,
|
||||||
|
):
|
||||||
|
state_dict = {}
|
||||||
|
|
||||||
|
def update_sd(prefix, sd):
|
||||||
|
for k, v in sd.items():
|
||||||
|
key = prefix + k
|
||||||
|
if save_dtype is not None:
|
||||||
|
v = v.detach().clone().to("cpu").to(save_dtype)
|
||||||
|
state_dict[key] = v
|
||||||
|
|
||||||
|
# Convert the UNet model
|
||||||
|
update_sd("model.diffusion_model.", unet.state_dict())
|
||||||
|
|
||||||
|
# Convert the text encoders
|
||||||
|
update_sd("conditioner.embedders.0.transformer.", text_encoder1.state_dict())
|
||||||
|
|
||||||
|
text_enc2_dict = convert_text_encoder_2_state_dict_to_sdxl(text_encoder2.state_dict(), logit_scale)
|
||||||
|
update_sd("conditioner.embedders.1.model.", text_enc2_dict)
|
||||||
|
|
||||||
|
# Convert the VAE
|
||||||
|
vae_dict = model_util.convert_vae_state_dict(vae.state_dict())
|
||||||
|
update_sd("first_stage_model.", vae_dict)
|
||||||
|
|
||||||
|
# Put together new checkpoint
|
||||||
|
key_count = len(state_dict.keys())
|
||||||
|
new_ckpt = {"state_dict": state_dict}
|
||||||
|
|
||||||
|
# epoch and global_step are sometimes not int
|
||||||
|
if ckpt_info is not None:
|
||||||
|
epochs += ckpt_info[0]
|
||||||
|
steps += ckpt_info[1]
|
||||||
|
|
||||||
|
new_ckpt["epoch"] = epochs
|
||||||
|
new_ckpt["global_step"] = steps
|
||||||
|
|
||||||
|
if model_util.is_safetensors(output_file):
|
||||||
|
save_file(state_dict, output_file, metadata)
|
||||||
|
else:
|
||||||
|
torch.save(new_ckpt, output_file)
|
||||||
|
|
||||||
|
return key_count
|
||||||
|
|
||||||
|
|
||||||
|
def save_diffusers_checkpoint(
|
||||||
|
output_dir, text_encoder1, text_encoder2, unet, pretrained_model_name_or_path, vae=None, use_safetensors=False, save_dtype=None
|
||||||
|
):
|
||||||
|
from diffusers import StableDiffusionXLPipeline
|
||||||
|
|
||||||
|
# convert U-Net
|
||||||
|
unet_sd = unet.state_dict()
|
||||||
|
du_unet_sd = convert_sdxl_unet_state_dict_to_diffusers(unet_sd)
|
||||||
|
|
||||||
|
diffusers_unet = UNet2DConditionModel(**DIFFUSERS_SDXL_UNET_CONFIG)
|
||||||
|
if save_dtype is not None:
|
||||||
|
diffusers_unet.to(save_dtype)
|
||||||
|
diffusers_unet.load_state_dict(du_unet_sd)
|
||||||
|
|
||||||
|
# create pipeline to save
|
||||||
|
if pretrained_model_name_or_path is None:
|
||||||
|
pretrained_model_name_or_path = DIFFUSERS_REF_MODEL_ID_SDXL
|
||||||
|
|
||||||
|
scheduler = EulerDiscreteScheduler.from_pretrained(pretrained_model_name_or_path, subfolder="scheduler")
|
||||||
|
tokenizer1 = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer")
|
||||||
|
tokenizer2 = CLIPTokenizer.from_pretrained(pretrained_model_name_or_path, subfolder="tokenizer_2")
|
||||||
|
if vae is None:
|
||||||
|
vae = AutoencoderKL.from_pretrained(pretrained_model_name_or_path, subfolder="vae")
|
||||||
|
|
||||||
|
# prevent local path from being saved
|
||||||
|
def remove_name_or_path(model):
|
||||||
|
if hasattr(model, "config"):
|
||||||
|
model.config._name_or_path = None
|
||||||
|
model.config._name_or_path = None
|
||||||
|
|
||||||
|
remove_name_or_path(diffusers_unet)
|
||||||
|
remove_name_or_path(text_encoder1)
|
||||||
|
remove_name_or_path(text_encoder2)
|
||||||
|
remove_name_or_path(scheduler)
|
||||||
|
remove_name_or_path(tokenizer1)
|
||||||
|
remove_name_or_path(tokenizer2)
|
||||||
|
remove_name_or_path(vae)
|
||||||
|
|
||||||
|
pipeline = StableDiffusionXLPipeline(
|
||||||
|
unet=diffusers_unet,
|
||||||
|
text_encoder=text_encoder1,
|
||||||
|
text_encoder_2=text_encoder2,
|
||||||
|
vae=vae,
|
||||||
|
scheduler=scheduler,
|
||||||
|
tokenizer=tokenizer1,
|
||||||
|
tokenizer_2=tokenizer2,
|
||||||
|
)
|
||||||
|
if save_dtype is not None:
|
||||||
|
pipeline.to(None, save_dtype)
|
||||||
|
pipeline.save_pretrained(output_dir, safe_serialization=use_safetensors)
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,381 @@
|
|||||||
|
import argparse
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from library.device_utils import init_ipex, clean_memory_on_device
|
||||||
|
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
from accelerate import init_empty_weights
|
||||||
|
from tqdm import tqdm
|
||||||
|
from transformers import CLIPTokenizer
|
||||||
|
from library import model_util, sdxl_model_util, train_util, sdxl_original_unet
|
||||||
|
from library.sdxl_lpw_stable_diffusion import SdxlStableDiffusionLongPromptWeightingPipeline
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
TOKENIZER1_PATH = "openai/clip-vit-large-patch14"
|
||||||
|
TOKENIZER2_PATH = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k"
|
||||||
|
|
||||||
|
# DEFAULT_NOISE_OFFSET = 0.0357
|
||||||
|
|
||||||
|
|
||||||
|
def load_target_model(args, accelerator, model_version: str, weight_dtype):
|
||||||
|
model_dtype = match_mixed_precision(args, weight_dtype) # prepare fp16/bf16
|
||||||
|
for pi in range(accelerator.state.num_processes):
|
||||||
|
if pi == accelerator.state.local_process_index:
|
||||||
|
logger.info(f"loading model for process {accelerator.state.local_process_index}/{accelerator.state.num_processes}")
|
||||||
|
|
||||||
|
(
|
||||||
|
load_stable_diffusion_format,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
vae,
|
||||||
|
unet,
|
||||||
|
logit_scale,
|
||||||
|
ckpt_info,
|
||||||
|
) = _load_target_model(
|
||||||
|
args.pretrained_model_name_or_path,
|
||||||
|
args.vae,
|
||||||
|
model_version,
|
||||||
|
weight_dtype,
|
||||||
|
accelerator.device if args.lowram else "cpu",
|
||||||
|
model_dtype,
|
||||||
|
args.disable_mmap_load_safetensors,
|
||||||
|
)
|
||||||
|
|
||||||
|
# work on low-ram device
|
||||||
|
if args.lowram:
|
||||||
|
text_encoder1.to(accelerator.device)
|
||||||
|
text_encoder2.to(accelerator.device)
|
||||||
|
unet.to(accelerator.device)
|
||||||
|
vae.to(accelerator.device)
|
||||||
|
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
|
return load_stable_diffusion_format, text_encoder1, text_encoder2, vae, unet, logit_scale, ckpt_info
|
||||||
|
|
||||||
|
|
||||||
|
def _load_target_model(
|
||||||
|
name_or_path: str, vae_path: Optional[str], model_version: str, weight_dtype, device="cpu", model_dtype=None, disable_mmap=False
|
||||||
|
):
|
||||||
|
# model_dtype only work with full fp16/bf16
|
||||||
|
name_or_path = os.readlink(name_or_path) if os.path.islink(name_or_path) else name_or_path
|
||||||
|
load_stable_diffusion_format = os.path.isfile(name_or_path) # determine SD or Diffusers
|
||||||
|
|
||||||
|
if load_stable_diffusion_format:
|
||||||
|
logger.info(f"load StableDiffusion checkpoint: {name_or_path}")
|
||||||
|
(
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
vae,
|
||||||
|
unet,
|
||||||
|
logit_scale,
|
||||||
|
ckpt_info,
|
||||||
|
) = sdxl_model_util.load_models_from_sdxl_checkpoint(model_version, name_or_path, device, model_dtype, disable_mmap)
|
||||||
|
else:
|
||||||
|
# Diffusers model is loaded to CPU
|
||||||
|
from diffusers import StableDiffusionXLPipeline
|
||||||
|
|
||||||
|
variant = "fp16" if weight_dtype == torch.float16 else None
|
||||||
|
logger.info(f"load Diffusers pretrained models: {name_or_path}, variant={variant}")
|
||||||
|
try:
|
||||||
|
try:
|
||||||
|
pipe = StableDiffusionXLPipeline.from_pretrained(
|
||||||
|
name_or_path, torch_dtype=model_dtype, variant=variant, tokenizer=None
|
||||||
|
)
|
||||||
|
except EnvironmentError as ex:
|
||||||
|
if variant is not None:
|
||||||
|
logger.info("try to load fp32 model")
|
||||||
|
pipe = StableDiffusionXLPipeline.from_pretrained(name_or_path, variant=None, tokenizer=None)
|
||||||
|
else:
|
||||||
|
raise ex
|
||||||
|
except EnvironmentError as ex:
|
||||||
|
logger.error(
|
||||||
|
f"model is not found as a file or in Hugging Face, perhaps file name is wrong? / 指定したモデル名のファイル、またはHugging Faceのモデルが見つかりません。ファイル名が誤っているかもしれません: {name_or_path}"
|
||||||
|
)
|
||||||
|
raise ex
|
||||||
|
|
||||||
|
text_encoder1 = pipe.text_encoder
|
||||||
|
text_encoder2 = pipe.text_encoder_2
|
||||||
|
|
||||||
|
# convert to fp32 for cache text_encoders outputs
|
||||||
|
if text_encoder1.dtype != torch.float32:
|
||||||
|
text_encoder1 = text_encoder1.to(dtype=torch.float32)
|
||||||
|
if text_encoder2.dtype != torch.float32:
|
||||||
|
text_encoder2 = text_encoder2.to(dtype=torch.float32)
|
||||||
|
|
||||||
|
vae = pipe.vae
|
||||||
|
unet = pipe.unet
|
||||||
|
del pipe
|
||||||
|
|
||||||
|
# Diffusers U-Net to original U-Net
|
||||||
|
state_dict = sdxl_model_util.convert_diffusers_unet_state_dict_to_sdxl(unet.state_dict())
|
||||||
|
with init_empty_weights():
|
||||||
|
unet = sdxl_original_unet.SdxlUNet2DConditionModel() # overwrite unet
|
||||||
|
sdxl_model_util._load_state_dict_on_device(unet, state_dict, device=device, dtype=model_dtype)
|
||||||
|
logger.info("U-Net converted to original U-Net")
|
||||||
|
|
||||||
|
logit_scale = None
|
||||||
|
ckpt_info = None
|
||||||
|
|
||||||
|
# VAEを読み込む
|
||||||
|
if vae_path is not None:
|
||||||
|
vae = model_util.load_vae(vae_path, weight_dtype)
|
||||||
|
logger.info("additional VAE loaded")
|
||||||
|
|
||||||
|
return load_stable_diffusion_format, text_encoder1, text_encoder2, vae, unet, logit_scale, ckpt_info
|
||||||
|
|
||||||
|
|
||||||
|
def load_tokenizers(args: argparse.Namespace):
|
||||||
|
logger.info("prepare tokenizers")
|
||||||
|
|
||||||
|
original_paths = [TOKENIZER1_PATH, TOKENIZER2_PATH]
|
||||||
|
tokeniers = []
|
||||||
|
for i, original_path in enumerate(original_paths):
|
||||||
|
tokenizer: CLIPTokenizer = None
|
||||||
|
if args.tokenizer_cache_dir:
|
||||||
|
local_tokenizer_path = os.path.join(args.tokenizer_cache_dir, original_path.replace("/", "_"))
|
||||||
|
if os.path.exists(local_tokenizer_path):
|
||||||
|
logger.info(f"load tokenizer from cache: {local_tokenizer_path}")
|
||||||
|
tokenizer = CLIPTokenizer.from_pretrained(local_tokenizer_path)
|
||||||
|
|
||||||
|
if tokenizer is None:
|
||||||
|
tokenizer = CLIPTokenizer.from_pretrained(original_path)
|
||||||
|
|
||||||
|
if args.tokenizer_cache_dir and not os.path.exists(local_tokenizer_path):
|
||||||
|
logger.info(f"save Tokenizer to cache: {local_tokenizer_path}")
|
||||||
|
tokenizer.save_pretrained(local_tokenizer_path)
|
||||||
|
|
||||||
|
if i == 1:
|
||||||
|
tokenizer.pad_token_id = 0 # fix pad token id to make same as open clip tokenizer
|
||||||
|
|
||||||
|
tokeniers.append(tokenizer)
|
||||||
|
|
||||||
|
if hasattr(args, "max_token_length") and args.max_token_length is not None:
|
||||||
|
logger.info(f"update token length: {args.max_token_length}")
|
||||||
|
|
||||||
|
return tokeniers
|
||||||
|
|
||||||
|
|
||||||
|
def match_mixed_precision(args, weight_dtype):
|
||||||
|
if args.full_fp16:
|
||||||
|
assert (
|
||||||
|
weight_dtype == torch.float16
|
||||||
|
), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"
|
||||||
|
return weight_dtype
|
||||||
|
elif args.full_bf16:
|
||||||
|
assert (
|
||||||
|
weight_dtype == torch.bfloat16
|
||||||
|
), "full_bf16 requires mixed precision='bf16' / full_bf16を使う場合はmixed_precision='bf16'を指定してください。"
|
||||||
|
return weight_dtype
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def timestep_embedding(timesteps, dim, max_period=10000):
|
||||||
|
"""
|
||||||
|
Create sinusoidal timestep embeddings.
|
||||||
|
:param timesteps: a 1-D Tensor of N indices, one per batch element.
|
||||||
|
These may be fractional.
|
||||||
|
:param dim: the dimension of the output.
|
||||||
|
:param max_period: controls the minimum frequency of the embeddings.
|
||||||
|
:return: an [N x dim] Tensor of positional embeddings.
|
||||||
|
"""
|
||||||
|
half = dim // 2
|
||||||
|
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(
|
||||||
|
device=timesteps.device
|
||||||
|
)
|
||||||
|
args = timesteps[:, None].float() * freqs[None]
|
||||||
|
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||||
|
if dim % 2:
|
||||||
|
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||||
|
return embedding
|
||||||
|
|
||||||
|
|
||||||
|
def get_timestep_embedding(x, outdim):
|
||||||
|
assert len(x.shape) == 2
|
||||||
|
b, dims = x.shape[0], x.shape[1]
|
||||||
|
x = torch.flatten(x)
|
||||||
|
emb = timestep_embedding(x, outdim)
|
||||||
|
emb = torch.reshape(emb, (b, dims * outdim))
|
||||||
|
return emb
|
||||||
|
|
||||||
|
|
||||||
|
def get_size_embeddings(orig_size, crop_size, target_size, device):
|
||||||
|
emb1 = get_timestep_embedding(orig_size, 256)
|
||||||
|
emb2 = get_timestep_embedding(crop_size, 256)
|
||||||
|
emb3 = get_timestep_embedding(target_size, 256)
|
||||||
|
vector = torch.cat([emb1, emb2, emb3], dim=1).to(device)
|
||||||
|
return vector
|
||||||
|
|
||||||
|
|
||||||
|
def save_sd_model_on_train_end(
|
||||||
|
args: argparse.Namespace,
|
||||||
|
src_path: str,
|
||||||
|
save_stable_diffusion_format: bool,
|
||||||
|
use_safetensors: bool,
|
||||||
|
save_dtype: torch.dtype,
|
||||||
|
epoch: int,
|
||||||
|
global_step: int,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
vae,
|
||||||
|
logit_scale,
|
||||||
|
ckpt_info,
|
||||||
|
):
|
||||||
|
def sd_saver(ckpt_file, epoch_no, global_step):
|
||||||
|
sai_metadata = train_util.get_sai_model_spec(None, args, True, False, False, is_stable_diffusion_ckpt=True)
|
||||||
|
sdxl_model_util.save_stable_diffusion_checkpoint(
|
||||||
|
ckpt_file,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
epoch_no,
|
||||||
|
global_step,
|
||||||
|
ckpt_info,
|
||||||
|
vae,
|
||||||
|
logit_scale,
|
||||||
|
sai_metadata,
|
||||||
|
save_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
def diffusers_saver(out_dir):
|
||||||
|
sdxl_model_util.save_diffusers_checkpoint(
|
||||||
|
out_dir,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
src_path,
|
||||||
|
vae,
|
||||||
|
use_safetensors=use_safetensors,
|
||||||
|
save_dtype=save_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
train_util.save_sd_model_on_train_end_common(
|
||||||
|
args, save_stable_diffusion_format, use_safetensors, epoch, global_step, sd_saver, diffusers_saver
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# epochとstepの保存、メタデータにepoch/stepが含まれ引数が同じになるため、統合している
|
||||||
|
# on_epoch_end: Trueならepoch終了時、Falseならstep経過時
|
||||||
|
def save_sd_model_on_epoch_end_or_stepwise(
|
||||||
|
args: argparse.Namespace,
|
||||||
|
on_epoch_end: bool,
|
||||||
|
accelerator,
|
||||||
|
src_path,
|
||||||
|
save_stable_diffusion_format: bool,
|
||||||
|
use_safetensors: bool,
|
||||||
|
save_dtype: torch.dtype,
|
||||||
|
epoch: int,
|
||||||
|
num_train_epochs: int,
|
||||||
|
global_step: int,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
vae,
|
||||||
|
logit_scale,
|
||||||
|
ckpt_info,
|
||||||
|
):
|
||||||
|
def sd_saver(ckpt_file, epoch_no, global_step):
|
||||||
|
sai_metadata = train_util.get_sai_model_spec(None, args, True, False, False, is_stable_diffusion_ckpt=True)
|
||||||
|
sdxl_model_util.save_stable_diffusion_checkpoint(
|
||||||
|
ckpt_file,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
epoch_no,
|
||||||
|
global_step,
|
||||||
|
ckpt_info,
|
||||||
|
vae,
|
||||||
|
logit_scale,
|
||||||
|
sai_metadata,
|
||||||
|
save_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
def diffusers_saver(out_dir):
|
||||||
|
sdxl_model_util.save_diffusers_checkpoint(
|
||||||
|
out_dir,
|
||||||
|
text_encoder1,
|
||||||
|
text_encoder2,
|
||||||
|
unet,
|
||||||
|
src_path,
|
||||||
|
vae,
|
||||||
|
use_safetensors=use_safetensors,
|
||||||
|
save_dtype=save_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise_common(
|
||||||
|
args,
|
||||||
|
on_epoch_end,
|
||||||
|
accelerator,
|
||||||
|
save_stable_diffusion_format,
|
||||||
|
use_safetensors,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
sd_saver,
|
||||||
|
diffusers_saver,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def add_sdxl_training_arguments(parser: argparse.ArgumentParser, support_text_encoder_caching: bool = True):
|
||||||
|
parser.add_argument(
|
||||||
|
"--cache_text_encoder_outputs", action="store_true", help="cache text encoder outputs / text encoderの出力をキャッシュする"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--cache_text_encoder_outputs_to_disk",
|
||||||
|
action="store_true",
|
||||||
|
help="cache text encoder outputs to disk / text encoderの出力をディスクにキャッシュする",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--disable_mmap_load_safetensors",
|
||||||
|
action="store_true",
|
||||||
|
help="disable mmap load for safetensors. Speed up model loading in WSL environment / safetensorsのmmapロードを無効にする。WSL環境等でモデル読み込みを高速化できる",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def verify_sdxl_training_args(args: argparse.Namespace, supportTextEncoderCaching: bool = True):
|
||||||
|
assert not args.v2, "v2 cannot be enabled in SDXL training / SDXL学習ではv2を有効にすることはできません"
|
||||||
|
if args.v_parameterization:
|
||||||
|
logger.warning("v_parameterization will be unexpected / SDXL学習ではv_parameterizationは想定外の動作になります")
|
||||||
|
|
||||||
|
if args.clip_skip is not None:
|
||||||
|
logger.warning("clip_skip will be unexpected / SDXL学習ではclip_skipは動作しません")
|
||||||
|
|
||||||
|
# if args.multires_noise_iterations:
|
||||||
|
# logger.info(
|
||||||
|
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET}, but noise_offset is disabled due to multires_noise_iterations / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されていますが、multires_noise_iterationsが有効になっているためnoise_offsetは無効になります"
|
||||||
|
# )
|
||||||
|
# else:
|
||||||
|
# if args.noise_offset is None:
|
||||||
|
# args.noise_offset = DEFAULT_NOISE_OFFSET
|
||||||
|
# elif args.noise_offset != DEFAULT_NOISE_OFFSET:
|
||||||
|
# logger.info(
|
||||||
|
# f"Warning: SDXL has been trained with noise_offset={DEFAULT_NOISE_OFFSET} / SDXLはnoise_offset={DEFAULT_NOISE_OFFSET}で学習されています"
|
||||||
|
# )
|
||||||
|
# logger.info(f"noise_offset is set to {args.noise_offset} / noise_offsetが{args.noise_offset}に設定されました")
|
||||||
|
|
||||||
|
assert (
|
||||||
|
not hasattr(args, "weighted_captions") or not args.weighted_captions
|
||||||
|
), "weighted_captions cannot be enabled in SDXL training currently / SDXL学習では今のところweighted_captionsを有効にすることはできません"
|
||||||
|
|
||||||
|
if supportTextEncoderCaching:
|
||||||
|
if args.cache_text_encoder_outputs_to_disk and not args.cache_text_encoder_outputs:
|
||||||
|
args.cache_text_encoder_outputs = True
|
||||||
|
logger.warning(
|
||||||
|
"cache_text_encoder_outputs is enabled because cache_text_encoder_outputs_to_disk is enabled / "
|
||||||
|
+ "cache_text_encoder_outputs_to_diskが有効になっているためcache_text_encoder_outputsが有効になりました"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def sample_images(*args, **kwargs):
|
||||||
|
return train_util.sample_images_common(SdxlStableDiffusionLongPromptWeightingPipeline, *args, **kwargs)
|
||||||
@@ -0,0 +1,682 @@
|
|||||||
|
# Modified from Diffusers to reduce VRAM usage
|
||||||
|
|
||||||
|
# Copyright 2022 The HuggingFace Team. All rights reserved.
|
||||||
|
#
|
||||||
|
# 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.
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers.models.modeling_utils import ModelMixin
|
||||||
|
from diffusers.models.unet_2d_blocks import UNetMidBlock2D, get_down_block, get_up_block
|
||||||
|
from diffusers.models.vae import DecoderOutput, DiagonalGaussianDistribution
|
||||||
|
from diffusers.models.autoencoder_kl import AutoencoderKLOutput
|
||||||
|
from .utils import setup_logging
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def slice_h(x, num_slices):
|
||||||
|
# slice with pad 1 both sides: to eliminate side effect of padding of conv2d
|
||||||
|
# Conv2dのpaddingの副作用を排除するために、両側にpad 1しながらHをスライスする
|
||||||
|
# NCHWでもNHWCでもどちらでも動く
|
||||||
|
size = (x.shape[2] + num_slices - 1) // num_slices
|
||||||
|
sliced = []
|
||||||
|
for i in range(num_slices):
|
||||||
|
if i == 0:
|
||||||
|
sliced.append(x[:, :, : size + 1, :])
|
||||||
|
else:
|
||||||
|
end = size * (i + 1) + 1
|
||||||
|
if x.shape[2] - end < 3: # if the last slice is too small, use the rest of the tensor 最後が細すぎるとconv2dできないので全部使う
|
||||||
|
end = x.shape[2]
|
||||||
|
sliced.append(x[:, :, size * i - 1 : end, :])
|
||||||
|
if end >= x.shape[2]:
|
||||||
|
break
|
||||||
|
return sliced
|
||||||
|
|
||||||
|
|
||||||
|
def cat_h(sliced):
|
||||||
|
# padding分を除いて結合する
|
||||||
|
cat = []
|
||||||
|
for i, x in enumerate(sliced):
|
||||||
|
if i == 0:
|
||||||
|
cat.append(x[:, :, :-1, :])
|
||||||
|
elif i == len(sliced) - 1:
|
||||||
|
cat.append(x[:, :, 1:, :])
|
||||||
|
else:
|
||||||
|
cat.append(x[:, :, 1:-1, :])
|
||||||
|
del x
|
||||||
|
x = torch.cat(cat, dim=2)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def resblock_forward(_self, num_slices, input_tensor, temb, **kwargs):
|
||||||
|
assert _self.upsample is None and _self.downsample is None
|
||||||
|
assert _self.norm1.num_groups == _self.norm2.num_groups
|
||||||
|
assert temb is None
|
||||||
|
|
||||||
|
# make sure norms are on cpu
|
||||||
|
org_device = input_tensor.device
|
||||||
|
cpu_device = torch.device("cpu")
|
||||||
|
_self.norm1.to(cpu_device)
|
||||||
|
_self.norm2.to(cpu_device)
|
||||||
|
|
||||||
|
# GroupNormがCPUでfp16で動かない対策
|
||||||
|
org_dtype = input_tensor.dtype
|
||||||
|
if org_dtype == torch.float16:
|
||||||
|
_self.norm1.to(torch.float32)
|
||||||
|
_self.norm2.to(torch.float32)
|
||||||
|
|
||||||
|
# すべてのテンソルをCPUに移動する
|
||||||
|
input_tensor = input_tensor.to(cpu_device)
|
||||||
|
hidden_states = input_tensor
|
||||||
|
|
||||||
|
# どうもこれは結果が異なるようだ……
|
||||||
|
# def sliced_norm1(norm, x):
|
||||||
|
# num_div = 4 if up_block_idx <= 2 else x.shape[1] // norm.num_groups
|
||||||
|
# sliced_tensor = torch.chunk(x, num_div, dim=1)
|
||||||
|
# sliced_weight = torch.chunk(norm.weight, num_div, dim=0)
|
||||||
|
# sliced_bias = torch.chunk(norm.bias, num_div, dim=0)
|
||||||
|
# logger.info(sliced_tensor[0].shape, num_div, sliced_weight[0].shape, sliced_bias[0].shape)
|
||||||
|
# normed_tensor = []
|
||||||
|
# for i in range(num_div):
|
||||||
|
# n = torch.group_norm(sliced_tensor[i], norm.num_groups, sliced_weight[i], sliced_bias[i], norm.eps)
|
||||||
|
# normed_tensor.append(n)
|
||||||
|
# del n
|
||||||
|
# x = torch.cat(normed_tensor, dim=1)
|
||||||
|
# return num_div, x
|
||||||
|
|
||||||
|
# normを分割すると結果が変わるので、ここだけは分割しない。GPUで計算するとVRAMが足りなくなるので、CPUで計算する。幸いCPUでもそこまで遅くない
|
||||||
|
if org_dtype == torch.float16:
|
||||||
|
hidden_states = hidden_states.to(torch.float32)
|
||||||
|
hidden_states = _self.norm1(hidden_states) # run on cpu
|
||||||
|
if org_dtype == torch.float16:
|
||||||
|
hidden_states = hidden_states.to(torch.float16)
|
||||||
|
|
||||||
|
sliced = slice_h(hidden_states, num_slices)
|
||||||
|
del hidden_states
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
# 計算する部分だけGPUに移動する、以下同様
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = _self.nonlinearity(x)
|
||||||
|
x = _self.conv1(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
hidden_states = cat_h(sliced)
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
if org_dtype == torch.float16:
|
||||||
|
hidden_states = hidden_states.to(torch.float32)
|
||||||
|
hidden_states = _self.norm2(hidden_states) # run on cpu
|
||||||
|
if org_dtype == torch.float16:
|
||||||
|
hidden_states = hidden_states.to(torch.float16)
|
||||||
|
|
||||||
|
sliced = slice_h(hidden_states, num_slices)
|
||||||
|
del hidden_states
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = _self.nonlinearity(x)
|
||||||
|
x = _self.dropout(x)
|
||||||
|
x = _self.conv2(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
hidden_states = cat_h(sliced)
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
# make shortcut
|
||||||
|
if _self.conv_shortcut is not None:
|
||||||
|
sliced = list(torch.chunk(input_tensor, num_slices, dim=2)) # no padding in conv_shortcut パディングがないので普通にスライスする
|
||||||
|
del input_tensor
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = _self.conv_shortcut(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
input_tensor = torch.cat(sliced, dim=2)
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
output_tensor = (input_tensor + hidden_states) / _self.output_scale_factor
|
||||||
|
|
||||||
|
output_tensor = output_tensor.to(org_device) # 次のレイヤーがGPUで計算する
|
||||||
|
return output_tensor
|
||||||
|
|
||||||
|
|
||||||
|
class SlicingEncoder(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels=3,
|
||||||
|
out_channels=3,
|
||||||
|
down_block_types=("DownEncoderBlock2D",),
|
||||||
|
block_out_channels=(64,),
|
||||||
|
layers_per_block=2,
|
||||||
|
norm_num_groups=32,
|
||||||
|
act_fn="silu",
|
||||||
|
double_z=True,
|
||||||
|
num_slices=2,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.layers_per_block = layers_per_block
|
||||||
|
|
||||||
|
self.conv_in = torch.nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
|
||||||
|
|
||||||
|
self.mid_block = None
|
||||||
|
self.down_blocks = nn.ModuleList([])
|
||||||
|
|
||||||
|
# down
|
||||||
|
output_channel = block_out_channels[0]
|
||||||
|
for i, down_block_type in enumerate(down_block_types):
|
||||||
|
input_channel = output_channel
|
||||||
|
output_channel = block_out_channels[i]
|
||||||
|
is_final_block = i == len(block_out_channels) - 1
|
||||||
|
|
||||||
|
down_block = get_down_block(
|
||||||
|
down_block_type,
|
||||||
|
num_layers=self.layers_per_block,
|
||||||
|
in_channels=input_channel,
|
||||||
|
out_channels=output_channel,
|
||||||
|
add_downsample=not is_final_block,
|
||||||
|
resnet_eps=1e-6,
|
||||||
|
downsample_padding=0,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
attention_head_dim=output_channel,
|
||||||
|
temb_channels=None,
|
||||||
|
)
|
||||||
|
self.down_blocks.append(down_block)
|
||||||
|
|
||||||
|
# mid
|
||||||
|
self.mid_block = UNetMidBlock2D(
|
||||||
|
in_channels=block_out_channels[-1],
|
||||||
|
resnet_eps=1e-6,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
output_scale_factor=1,
|
||||||
|
resnet_time_scale_shift="default",
|
||||||
|
attention_head_dim=block_out_channels[-1],
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
temb_channels=None,
|
||||||
|
)
|
||||||
|
self.mid_block.attentions[0].set_use_memory_efficient_attention_xformers(True) # とりあえずDiffusersのxformersを使う
|
||||||
|
|
||||||
|
# out
|
||||||
|
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||||
|
self.conv_act = nn.SiLU()
|
||||||
|
|
||||||
|
conv_out_channels = 2 * out_channels if double_z else out_channels
|
||||||
|
self.conv_out = nn.Conv2d(block_out_channels[-1], conv_out_channels, 3, padding=1)
|
||||||
|
|
||||||
|
# replace forward of ResBlocks
|
||||||
|
def wrapper(func, module, num_slices):
|
||||||
|
def forward(*args, **kwargs):
|
||||||
|
return func(module, num_slices, *args, **kwargs)
|
||||||
|
|
||||||
|
return forward
|
||||||
|
|
||||||
|
self.num_slices = num_slices
|
||||||
|
div = num_slices / (2 ** (len(self.down_blocks) - 1)) # 深い層はそこまで分割しなくていいので適宜減らす
|
||||||
|
# logger.info(f"initial divisor: {div}")
|
||||||
|
if div >= 2:
|
||||||
|
div = int(div)
|
||||||
|
for resnet in self.mid_block.resnets:
|
||||||
|
resnet.forward = wrapper(resblock_forward, resnet, div)
|
||||||
|
# midblock doesn't have downsample
|
||||||
|
|
||||||
|
for i, down_block in enumerate(self.down_blocks[::-1]):
|
||||||
|
if div >= 2:
|
||||||
|
div = int(div)
|
||||||
|
# logger.info(f"down block: {i} divisor: {div}")
|
||||||
|
for resnet in down_block.resnets:
|
||||||
|
resnet.forward = wrapper(resblock_forward, resnet, div)
|
||||||
|
if down_block.downsamplers is not None:
|
||||||
|
# logger.info("has downsample")
|
||||||
|
for downsample in down_block.downsamplers:
|
||||||
|
downsample.forward = wrapper(self.downsample_forward, downsample, div * 2)
|
||||||
|
div *= 2
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
sample = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
org_device = sample.device
|
||||||
|
cpu_device = torch.device("cpu")
|
||||||
|
|
||||||
|
# sample = self.conv_in(sample)
|
||||||
|
sample = sample.to(cpu_device)
|
||||||
|
sliced = slice_h(sample, self.num_slices)
|
||||||
|
del sample
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = self.conv_in(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
sample = cat_h(sliced)
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
sample = sample.to(org_device)
|
||||||
|
|
||||||
|
# down
|
||||||
|
for down_block in self.down_blocks:
|
||||||
|
sample = down_block(sample)
|
||||||
|
|
||||||
|
# middle
|
||||||
|
sample = self.mid_block(sample)
|
||||||
|
|
||||||
|
# post-process
|
||||||
|
# ここも省メモリ化したいが、恐らくそこまでメモリを食わないので省略
|
||||||
|
sample = self.conv_norm_out(sample)
|
||||||
|
sample = self.conv_act(sample)
|
||||||
|
sample = self.conv_out(sample)
|
||||||
|
|
||||||
|
return sample
|
||||||
|
|
||||||
|
def downsample_forward(self, _self, num_slices, hidden_states):
|
||||||
|
assert hidden_states.shape[1] == _self.channels
|
||||||
|
assert _self.use_conv and _self.padding == 0
|
||||||
|
logger.info(f"downsample forward {num_slices} {hidden_states.shape}")
|
||||||
|
|
||||||
|
org_device = hidden_states.device
|
||||||
|
cpu_device = torch.device("cpu")
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(cpu_device)
|
||||||
|
pad = (0, 1, 0, 1)
|
||||||
|
hidden_states = torch.nn.functional.pad(hidden_states, pad, mode="constant", value=0)
|
||||||
|
|
||||||
|
# slice with even number because of stride 2
|
||||||
|
# strideが2なので偶数でスライスする
|
||||||
|
# slice with pad 1 both sides: to eliminate side effect of padding of conv2d
|
||||||
|
size = (hidden_states.shape[2] + num_slices - 1) // num_slices
|
||||||
|
size = size + 1 if size % 2 == 1 else size
|
||||||
|
|
||||||
|
sliced = []
|
||||||
|
for i in range(num_slices):
|
||||||
|
if i == 0:
|
||||||
|
sliced.append(hidden_states[:, :, : size + 1, :])
|
||||||
|
else:
|
||||||
|
end = size * (i + 1) + 1
|
||||||
|
if hidden_states.shape[2] - end < 4: # if the last slice is too small, use the rest of the tensor
|
||||||
|
end = hidden_states.shape[2]
|
||||||
|
sliced.append(hidden_states[:, :, size * i - 1 : end, :])
|
||||||
|
if end >= hidden_states.shape[2]:
|
||||||
|
break
|
||||||
|
del hidden_states
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = _self.conv(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
|
||||||
|
# ここだけ雰囲気が違うのはCopilotのせい
|
||||||
|
if i == 0:
|
||||||
|
hidden_states = x
|
||||||
|
else:
|
||||||
|
hidden_states = torch.cat([hidden_states, x], dim=2)
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(org_device)
|
||||||
|
# logger.info(f"downsample forward done {hidden_states.shape}")
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class SlicingDecoder(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels=3,
|
||||||
|
out_channels=3,
|
||||||
|
up_block_types=("UpDecoderBlock2D",),
|
||||||
|
block_out_channels=(64,),
|
||||||
|
layers_per_block=2,
|
||||||
|
norm_num_groups=32,
|
||||||
|
act_fn="silu",
|
||||||
|
num_slices=2,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.layers_per_block = layers_per_block
|
||||||
|
|
||||||
|
self.conv_in = nn.Conv2d(in_channels, block_out_channels[-1], kernel_size=3, stride=1, padding=1)
|
||||||
|
|
||||||
|
self.mid_block = None
|
||||||
|
self.up_blocks = nn.ModuleList([])
|
||||||
|
|
||||||
|
# mid
|
||||||
|
self.mid_block = UNetMidBlock2D(
|
||||||
|
in_channels=block_out_channels[-1],
|
||||||
|
resnet_eps=1e-6,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
output_scale_factor=1,
|
||||||
|
resnet_time_scale_shift="default",
|
||||||
|
attention_head_dim=block_out_channels[-1],
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
temb_channels=None,
|
||||||
|
)
|
||||||
|
self.mid_block.attentions[0].set_use_memory_efficient_attention_xformers(True) # とりあえずDiffusersのxformersを使う
|
||||||
|
|
||||||
|
# up
|
||||||
|
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||||
|
output_channel = reversed_block_out_channels[0]
|
||||||
|
for i, up_block_type in enumerate(up_block_types):
|
||||||
|
prev_output_channel = output_channel
|
||||||
|
output_channel = reversed_block_out_channels[i]
|
||||||
|
|
||||||
|
is_final_block = i == len(block_out_channels) - 1
|
||||||
|
|
||||||
|
up_block = get_up_block(
|
||||||
|
up_block_type,
|
||||||
|
num_layers=self.layers_per_block + 1,
|
||||||
|
in_channels=prev_output_channel,
|
||||||
|
out_channels=output_channel,
|
||||||
|
prev_output_channel=None,
|
||||||
|
add_upsample=not is_final_block,
|
||||||
|
resnet_eps=1e-6,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
attention_head_dim=output_channel,
|
||||||
|
temb_channels=None,
|
||||||
|
)
|
||||||
|
self.up_blocks.append(up_block)
|
||||||
|
prev_output_channel = output_channel
|
||||||
|
|
||||||
|
# out
|
||||||
|
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
|
||||||
|
self.conv_act = nn.SiLU()
|
||||||
|
self.conv_out = nn.Conv2d(block_out_channels[0], out_channels, 3, padding=1)
|
||||||
|
|
||||||
|
# replace forward of ResBlocks
|
||||||
|
def wrapper(func, module, num_slices):
|
||||||
|
def forward(*args, **kwargs):
|
||||||
|
return func(module, num_slices, *args, **kwargs)
|
||||||
|
|
||||||
|
return forward
|
||||||
|
|
||||||
|
self.num_slices = num_slices
|
||||||
|
div = num_slices / (2 ** (len(self.up_blocks) - 1))
|
||||||
|
logger.info(f"initial divisor: {div}")
|
||||||
|
if div >= 2:
|
||||||
|
div = int(div)
|
||||||
|
for resnet in self.mid_block.resnets:
|
||||||
|
resnet.forward = wrapper(resblock_forward, resnet, div)
|
||||||
|
# midblock doesn't have upsample
|
||||||
|
|
||||||
|
for i, up_block in enumerate(self.up_blocks):
|
||||||
|
if div >= 2:
|
||||||
|
div = int(div)
|
||||||
|
# logger.info(f"up block: {i} divisor: {div}")
|
||||||
|
for resnet in up_block.resnets:
|
||||||
|
resnet.forward = wrapper(resblock_forward, resnet, div)
|
||||||
|
if up_block.upsamplers is not None:
|
||||||
|
# logger.info("has upsample")
|
||||||
|
for upsample in up_block.upsamplers:
|
||||||
|
upsample.forward = wrapper(self.upsample_forward, upsample, div * 2)
|
||||||
|
div *= 2
|
||||||
|
|
||||||
|
def forward(self, z):
|
||||||
|
sample = z
|
||||||
|
del z
|
||||||
|
sample = self.conv_in(sample)
|
||||||
|
|
||||||
|
# middle
|
||||||
|
sample = self.mid_block(sample)
|
||||||
|
|
||||||
|
# up
|
||||||
|
for i, up_block in enumerate(self.up_blocks):
|
||||||
|
sample = up_block(sample)
|
||||||
|
|
||||||
|
# post-process
|
||||||
|
sample = self.conv_norm_out(sample)
|
||||||
|
sample = self.conv_act(sample)
|
||||||
|
|
||||||
|
# conv_out with slicing because of VRAM usage
|
||||||
|
# conv_outはとてもVRAM使うのでスライスして対応
|
||||||
|
org_device = sample.device
|
||||||
|
cpu_device = torch.device("cpu")
|
||||||
|
sample = sample.to(cpu_device)
|
||||||
|
|
||||||
|
sliced = slice_h(sample, self.num_slices)
|
||||||
|
del sample
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
x = self.conv_out(x)
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
sample = cat_h(sliced)
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
sample = sample.to(org_device)
|
||||||
|
return sample
|
||||||
|
|
||||||
|
def upsample_forward(self, _self, num_slices, hidden_states, output_size=None):
|
||||||
|
assert hidden_states.shape[1] == _self.channels
|
||||||
|
assert _self.use_conv_transpose == False and _self.use_conv
|
||||||
|
|
||||||
|
org_dtype = hidden_states.dtype
|
||||||
|
org_device = hidden_states.device
|
||||||
|
cpu_device = torch.device("cpu")
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(cpu_device)
|
||||||
|
sliced = slice_h(hidden_states, num_slices)
|
||||||
|
del hidden_states
|
||||||
|
|
||||||
|
for i in range(len(sliced)):
|
||||||
|
x = sliced[i]
|
||||||
|
sliced[i] = None
|
||||||
|
|
||||||
|
x = x.to(org_device)
|
||||||
|
|
||||||
|
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||||
|
# TODO(Suraj): Remove this cast once the issue is fixed in PyTorch
|
||||||
|
# https://github.com/pytorch/pytorch/issues/86679
|
||||||
|
# PyTorch 2で直らないかね……
|
||||||
|
if org_dtype == torch.bfloat16:
|
||||||
|
x = x.to(torch.float32)
|
||||||
|
|
||||||
|
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
|
||||||
|
|
||||||
|
if org_dtype == torch.bfloat16:
|
||||||
|
x = x.to(org_dtype)
|
||||||
|
|
||||||
|
x = _self.conv(x)
|
||||||
|
|
||||||
|
# upsampleされてるのでpadは2になる
|
||||||
|
if i == 0:
|
||||||
|
x = x[:, :, :-2, :]
|
||||||
|
elif i == num_slices - 1:
|
||||||
|
x = x[:, :, 2:, :]
|
||||||
|
else:
|
||||||
|
x = x[:, :, 2:-2, :]
|
||||||
|
|
||||||
|
x = x.to(cpu_device)
|
||||||
|
sliced[i] = x
|
||||||
|
del x
|
||||||
|
|
||||||
|
hidden_states = torch.cat(sliced, dim=2)
|
||||||
|
# logger.info(f"us hidden_states {hidden_states.shape}")
|
||||||
|
del sliced
|
||||||
|
|
||||||
|
hidden_states = hidden_states.to(org_device)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class SlicingAutoencoderKL(ModelMixin, ConfigMixin):
|
||||||
|
r"""Variational Autoencoder (VAE) model with KL loss from the paper Auto-Encoding Variational Bayes by Diederik P. Kingma
|
||||||
|
and Max Welling.
|
||||||
|
|
||||||
|
This model inherits from [`ModelMixin`]. Check the superclass documentation for the generic methods the library
|
||||||
|
implements for all the model (such as downloading or saving, etc.)
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
in_channels (int, *optional*, defaults to 3): Number of channels in the input image.
|
||||||
|
out_channels (int, *optional*, defaults to 3): Number of channels in the output.
|
||||||
|
down_block_types (`Tuple[str]`, *optional*, defaults to :
|
||||||
|
obj:`("DownEncoderBlock2D",)`): Tuple of downsample block types.
|
||||||
|
up_block_types (`Tuple[str]`, *optional*, defaults to :
|
||||||
|
obj:`("UpDecoderBlock2D",)`): Tuple of upsample block types.
|
||||||
|
block_out_channels (`Tuple[int]`, *optional*, defaults to :
|
||||||
|
obj:`(64,)`): Tuple of block output channels.
|
||||||
|
act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use.
|
||||||
|
latent_channels (`int`, *optional*, defaults to `4`): Number of channels in the latent space.
|
||||||
|
sample_size (`int`, *optional*, defaults to `32`): TODO
|
||||||
|
"""
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int = 3,
|
||||||
|
out_channels: int = 3,
|
||||||
|
down_block_types: Tuple[str] = ("DownEncoderBlock2D",),
|
||||||
|
up_block_types: Tuple[str] = ("UpDecoderBlock2D",),
|
||||||
|
block_out_channels: Tuple[int] = (64,),
|
||||||
|
layers_per_block: int = 1,
|
||||||
|
act_fn: str = "silu",
|
||||||
|
latent_channels: int = 4,
|
||||||
|
norm_num_groups: int = 32,
|
||||||
|
sample_size: int = 32,
|
||||||
|
num_slices: int = 16,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
# pass init params to Encoder
|
||||||
|
self.encoder = SlicingEncoder(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=latent_channels,
|
||||||
|
down_block_types=down_block_types,
|
||||||
|
block_out_channels=block_out_channels,
|
||||||
|
layers_per_block=layers_per_block,
|
||||||
|
act_fn=act_fn,
|
||||||
|
norm_num_groups=norm_num_groups,
|
||||||
|
double_z=True,
|
||||||
|
num_slices=num_slices,
|
||||||
|
)
|
||||||
|
|
||||||
|
# pass init params to Decoder
|
||||||
|
self.decoder = SlicingDecoder(
|
||||||
|
in_channels=latent_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
up_block_types=up_block_types,
|
||||||
|
block_out_channels=block_out_channels,
|
||||||
|
layers_per_block=layers_per_block,
|
||||||
|
norm_num_groups=norm_num_groups,
|
||||||
|
act_fn=act_fn,
|
||||||
|
num_slices=num_slices,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.quant_conv = torch.nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1)
|
||||||
|
self.post_quant_conv = torch.nn.Conv2d(latent_channels, latent_channels, 1)
|
||||||
|
self.use_slicing = False
|
||||||
|
|
||||||
|
def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
|
||||||
|
h = self.encoder(x)
|
||||||
|
moments = self.quant_conv(h)
|
||||||
|
posterior = DiagonalGaussianDistribution(moments)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (posterior,)
|
||||||
|
|
||||||
|
return AutoencoderKLOutput(latent_dist=posterior)
|
||||||
|
|
||||||
|
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||||
|
z = self.post_quant_conv(z)
|
||||||
|
dec = self.decoder(z)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (dec,)
|
||||||
|
|
||||||
|
return DecoderOutput(sample=dec)
|
||||||
|
|
||||||
|
# これはバッチ方向のスライシング 紛らわしい
|
||||||
|
def enable_slicing(self):
|
||||||
|
r"""
|
||||||
|
Enable sliced VAE decoding.
|
||||||
|
|
||||||
|
When this option is enabled, the VAE will split the input tensor in slices to compute decoding in several
|
||||||
|
steps. This is useful to save some memory and allow larger batch sizes.
|
||||||
|
"""
|
||||||
|
self.use_slicing = True
|
||||||
|
|
||||||
|
def disable_slicing(self):
|
||||||
|
r"""
|
||||||
|
Disable sliced VAE decoding. If `enable_slicing` was previously invoked, this method will go back to computing
|
||||||
|
decoding in one step.
|
||||||
|
"""
|
||||||
|
self.use_slicing = False
|
||||||
|
|
||||||
|
def decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||||
|
if self.use_slicing and z.shape[0] > 1:
|
||||||
|
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
|
||||||
|
decoded = torch.cat(decoded_slices)
|
||||||
|
else:
|
||||||
|
decoded = self._decode(z).sample
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (decoded,)
|
||||||
|
|
||||||
|
return DecoderOutput(sample=decoded)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
sample_posterior: bool = False,
|
||||||
|
return_dict: bool = True,
|
||||||
|
generator: Optional[torch.Generator] = None,
|
||||||
|
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||||
|
r"""
|
||||||
|
Args:
|
||||||
|
sample (`torch.FloatTensor`): Input sample.
|
||||||
|
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to sample from the posterior.
|
||||||
|
return_dict (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
|
||||||
|
"""
|
||||||
|
x = sample
|
||||||
|
posterior = self.encode(x).latent_dist
|
||||||
|
if sample_posterior:
|
||||||
|
z = posterior.sample(generator=generator)
|
||||||
|
else:
|
||||||
|
z = posterior.mode()
|
||||||
|
dec = self.decode(z).sample
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (dec,)
|
||||||
|
|
||||||
|
return DecoderOutput(sample=dec)
|
||||||
@@ -0,0 +1,326 @@
|
|||||||
|
# base class for platform strategies. this file defines the interface for strategies
|
||||||
|
|
||||||
|
import os
|
||||||
|
from typing import Any, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from transformers import CLIPTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
# TODO remove circular import by moving ImageInfo to a separate file
|
||||||
|
# from library.train_util import ImageInfo
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenizeStrategy:
|
||||||
|
_strategy = None # strategy instance: actual strategy class
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def set_strategy(cls, strategy):
|
||||||
|
cls._strategy = strategy
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_strategy(cls):
|
||||||
|
return cls._strategy
|
||||||
|
|
||||||
|
def _load_tokenizer(
|
||||||
|
self, model_class: Any, model_id: str, subfolder: Optional[str] = None, tokenizer_cache_dir: Optional[str] = None
|
||||||
|
) -> Any:
|
||||||
|
tokenizer = None
|
||||||
|
if tokenizer_cache_dir:
|
||||||
|
local_tokenizer_path = os.path.join(tokenizer_cache_dir, model_id.replace("/", "_"))
|
||||||
|
if os.path.exists(local_tokenizer_path):
|
||||||
|
logger.info(f"load tokenizer from cache: {local_tokenizer_path}")
|
||||||
|
tokenizer = model_class.from_pretrained(local_tokenizer_path) # same for v1 and v2
|
||||||
|
|
||||||
|
if tokenizer is None:
|
||||||
|
tokenizer = model_class.from_pretrained(model_id, subfolder=subfolder)
|
||||||
|
|
||||||
|
if tokenizer_cache_dir and not os.path.exists(local_tokenizer_path):
|
||||||
|
logger.info(f"save Tokenizer to cache: {local_tokenizer_path}")
|
||||||
|
tokenizer.save_pretrained(local_tokenizer_path)
|
||||||
|
|
||||||
|
return tokenizer
|
||||||
|
|
||||||
|
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:
|
||||||
|
"""
|
||||||
|
for SD1.5/2.0/SDXL
|
||||||
|
TODO support batch input
|
||||||
|
"""
|
||||||
|
if max_length is None:
|
||||||
|
max_length = tokenizer.model_max_length - 2
|
||||||
|
|
||||||
|
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:
|
||||||
|
input_ids = input_ids.squeeze(0)
|
||||||
|
iids_list = []
|
||||||
|
if tokenizer.pad_token_id == tokenizer.eos_token_id:
|
||||||
|
# v1
|
||||||
|
# 77以上の時は "<BOS> .... <EOS> <EOS> <EOS>" でトータル227とかになっているので、"<BOS>...<EOS>"の三連に変換する
|
||||||
|
# 1111氏のやつは , で区切る、とかしているようだが とりあえず単純に
|
||||||
|
for i in range(1, max_length - tokenizer.model_max_length + 2, tokenizer.model_max_length - 2): # (1, 152, 75)
|
||||||
|
ids_chunk = (
|
||||||
|
input_ids[0].unsqueeze(0),
|
||||||
|
input_ids[i : i + tokenizer.model_max_length - 2],
|
||||||
|
input_ids[-1].unsqueeze(0),
|
||||||
|
)
|
||||||
|
ids_chunk = torch.cat(ids_chunk)
|
||||||
|
iids_list.append(ids_chunk)
|
||||||
|
else:
|
||||||
|
# v2 or SDXL
|
||||||
|
# 77以上の時は "<BOS> .... <EOS> <PAD> <PAD>..." でトータル227とかになっているので、"<BOS>...<EOS> <PAD> <PAD> ..."の三連に変換する
|
||||||
|
for i in range(1, max_length - tokenizer.model_max_length + 2, tokenizer.model_max_length - 2):
|
||||||
|
ids_chunk = (
|
||||||
|
input_ids[0].unsqueeze(0), # BOS
|
||||||
|
input_ids[i : i + tokenizer.model_max_length - 2],
|
||||||
|
input_ids[-1].unsqueeze(0),
|
||||||
|
) # PAD or EOS
|
||||||
|
ids_chunk = torch.cat(ids_chunk)
|
||||||
|
|
||||||
|
# 末尾が <EOS> <PAD> または <PAD> <PAD> の場合は、何もしなくてよい
|
||||||
|
# 末尾が x <PAD/EOS> の場合は末尾を <EOS> に変える(x <EOS> なら結果的に変化なし)
|
||||||
|
if ids_chunk[-2] != tokenizer.eos_token_id and ids_chunk[-2] != tokenizer.pad_token_id:
|
||||||
|
ids_chunk[-1] = tokenizer.eos_token_id
|
||||||
|
# 先頭が <BOS> <PAD> ... の場合は <BOS> <EOS> <PAD> ... に変える
|
||||||
|
if ids_chunk[1] == tokenizer.pad_token_id:
|
||||||
|
ids_chunk[1] = tokenizer.eos_token_id
|
||||||
|
|
||||||
|
iids_list.append(ids_chunk)
|
||||||
|
|
||||||
|
input_ids = torch.stack(iids_list) # 3,77
|
||||||
|
return input_ids
|
||||||
|
|
||||||
|
|
||||||
|
class TextEncodingStrategy:
|
||||||
|
_strategy = None # strategy instance: actual strategy class
|
||||||
|
|
||||||
|
@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
|
||||||
|
def get_strategy(cls) -> Optional["TextEncodingStrategy"]:
|
||||||
|
return cls._strategy
|
||||||
|
|
||||||
|
def encode_tokens(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], tokens: List[torch.Tensor]
|
||||||
|
) -> List[torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Encode tokens into embeddings and outputs.
|
||||||
|
:param tokens: list of token 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
|
||||||
|
) -> 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
|
||||||
|
|
||||||
|
@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
|
||||||
|
def get_strategy(cls) -> Optional["TextEncoderOutputsCachingStrategy"]:
|
||||||
|
return cls._strategy
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cache_to_disk(self):
|
||||||
|
return self._cache_to_disk
|
||||||
|
|
||||||
|
@property
|
||||||
|
def batch_size(self):
|
||||||
|
return self._batch_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_partial(self):
|
||||||
|
return self._is_partial
|
||||||
|
|
||||||
|
def get_outputs_npz_path(self, image_abs_path: str) -> str:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def load_outputs_npz(self, npz_path: str) -> List[np.ndarray]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def is_disk_cached_outputs_expected(self, npz_path: str) -> bool:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def cache_batch_outputs(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, batch: List
|
||||||
|
):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class LatentsCachingStrategy:
|
||||||
|
# TODO commonize utillity functions to this class, such as npz handling etc.
|
||||||
|
|
||||||
|
_strategy = None # strategy instance: actual strategy class
|
||||||
|
|
||||||
|
def __init__(self, cache_to_disk: bool, batch_size: int, skip_disk_cache_validity_check: bool) -> None:
|
||||||
|
self._cache_to_disk = cache_to_disk
|
||||||
|
self._batch_size = batch_size
|
||||||
|
self.skip_disk_cache_validity_check = skip_disk_cache_validity_check
|
||||||
|
|
||||||
|
@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
|
||||||
|
def get_strategy(cls) -> Optional["LatentsCachingStrategy"]:
|
||||||
|
return cls._strategy
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cache_to_disk(self):
|
||||||
|
return self._cache_to_disk
|
||||||
|
|
||||||
|
@property
|
||||||
|
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]]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def is_disk_cached_latents_expected(
|
||||||
|
self, bucket_reso: Tuple[int, int], npz_path: str, flip_aug: bool, alpha_mask: bool
|
||||||
|
) -> bool:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def cache_batch_latents(self, model: Any, batch: List, flip_aug: bool, alpha_mask: bool, random_crop: bool):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _default_is_disk_cached_latents_expected(
|
||||||
|
self, latents_stride: int, bucket_reso: Tuple[int, int], npz_path: str, flip_aug: bool, alpha_mask: bool
|
||||||
|
):
|
||||||
|
if not self.cache_to_disk:
|
||||||
|
return False
|
||||||
|
if not os.path.exists(npz_path):
|
||||||
|
return False
|
||||||
|
if self.skip_disk_cache_validity_check:
|
||||||
|
return True
|
||||||
|
|
||||||
|
expected_latents_size = (bucket_reso[1] // latents_stride, bucket_reso[0] // latents_stride) # bucket_reso is (W, H)
|
||||||
|
|
||||||
|
try:
|
||||||
|
npz = np.load(npz_path)
|
||||||
|
if npz["latents"].shape[1:3] != expected_latents_size:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if flip_aug:
|
||||||
|
if "latents_flipped" not in npz:
|
||||||
|
return False
|
||||||
|
if npz["latents_flipped"].shape[1:3] != expected_latents_size:
|
||||||
|
return False
|
||||||
|
|
||||||
|
if alpha_mask:
|
||||||
|
if "alpha_mask" not in npz:
|
||||||
|
return False
|
||||||
|
if npz["alpha_mask"].shape[0:2] != (bucket_reso[1], bucket_reso[0]):
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
if "alpha_mask" in npz:
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error loading file: {npz_path}")
|
||||||
|
raise e
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
# TODO remove circular dependency for ImageInfo
|
||||||
|
def _default_cache_batch_latents(
|
||||||
|
self, encode_by_vae, vae_device, vae_dtype, image_infos: List, flip_aug: bool, alpha_mask: bool, random_crop: bool
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Default implementation for cache_batch_latents. Image loading, VAE, flipping, alpha mask handling are common.
|
||||||
|
"""
|
||||||
|
from library import train_util # import here to avoid circular import
|
||||||
|
|
||||||
|
img_tensor, alpha_masks, original_sizes, crop_ltrbs = train_util.load_images_and_masks_for_caching(
|
||||||
|
image_infos, alpha_mask, random_crop
|
||||||
|
)
|
||||||
|
img_tensor = img_tensor.to(device=vae_device, dtype=vae_dtype)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
latents_tensors = encode_by_vae(img_tensor).to("cpu")
|
||||||
|
if flip_aug:
|
||||||
|
img_tensor = torch.flip(img_tensor, dims=[3])
|
||||||
|
with torch.no_grad():
|
||||||
|
flipped_latents = encode_by_vae(img_tensor).to("cpu")
|
||||||
|
else:
|
||||||
|
flipped_latents = [None] * len(latents_tensors)
|
||||||
|
|
||||||
|
# for info, latents, flipped_latent, alpha_mask in zip(image_infos, latents_tensors, flipped_latents, alpha_masks):
|
||||||
|
for i in range(len(image_infos)):
|
||||||
|
info = image_infos[i]
|
||||||
|
latents = latents_tensors[i]
|
||||||
|
flipped_latent = flipped_latents[i]
|
||||||
|
alpha_mask = alpha_masks[i]
|
||||||
|
original_size = original_sizes[i]
|
||||||
|
crop_ltrb = crop_ltrbs[i]
|
||||||
|
|
||||||
|
if self.cache_to_disk:
|
||||||
|
self.save_latents_to_disk(info.latents_npz, latents, original_size, crop_ltrb, flipped_latent, alpha_mask)
|
||||||
|
else:
|
||||||
|
info.latents_original_size = original_size
|
||||||
|
info.latents_crop_ltrb = crop_ltrb
|
||||||
|
info.latents = latents
|
||||||
|
if flip_aug:
|
||||||
|
info.latents_flipped = flipped_latent
|
||||||
|
info.alpha_mask = alpha_mask
|
||||||
|
|
||||||
|
def load_latents_from_disk(
|
||||||
|
self, npz_path: str
|
||||||
|
) -> Tuple[Optional[np.ndarray], Optional[List[int]], Optional[List[int]], Optional[np.ndarray], Optional[np.ndarray]]:
|
||||||
|
npz = np.load(npz_path)
|
||||||
|
if "latents" not in npz:
|
||||||
|
raise ValueError(f"error: npz is old format. please re-generate {npz_path}")
|
||||||
|
|
||||||
|
latents = npz["latents"]
|
||||||
|
original_size = npz["original_size"].tolist()
|
||||||
|
crop_ltrb = npz["crop_ltrb"].tolist()
|
||||||
|
flipped_latents = npz["latents_flipped"] if "latents_flipped" in npz else None
|
||||||
|
alpha_mask = npz["alpha_mask"] if "alpha_mask" in npz else None
|
||||||
|
return latents, original_size, crop_ltrb, flipped_latents, alpha_mask
|
||||||
|
|
||||||
|
def save_latents_to_disk(
|
||||||
|
self, npz_path, latents_tensor, original_size, crop_ltrb, flipped_latents_tensor=None, alpha_mask=None
|
||||||
|
):
|
||||||
|
kwargs = {}
|
||||||
|
if flipped_latents_tensor is not None:
|
||||||
|
kwargs["latents_flipped"] = flipped_latents_tensor.float().cpu().numpy()
|
||||||
|
if alpha_mask is not None:
|
||||||
|
kwargs["alpha_mask"] = alpha_mask.float().cpu().numpy()
|
||||||
|
np.savez(
|
||||||
|
npz_path,
|
||||||
|
latents=latents_tensor.float().cpu().numpy(),
|
||||||
|
original_size=np.array(original_size),
|
||||||
|
crop_ltrb=np.array(crop_ltrb),
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
@@ -0,0 +1,251 @@
|
|||||||
|
import os
|
||||||
|
import glob
|
||||||
|
from typing import Any, List, Optional, Tuple, Union
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
from transformers import CLIPTokenizer, T5TokenizerFast
|
||||||
|
|
||||||
|
from . import train_util
|
||||||
|
from .strategy_base import LatentsCachingStrategy, TextEncodingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
CLIP_L_TOKENIZER_ID = "openai/clip-vit-large-patch14"
|
||||||
|
T5_XXL_TOKENIZER_ID = "google/t5-v1_1-xxl"
|
||||||
|
|
||||||
|
|
||||||
|
class FluxTokenizeStrategy(TokenizeStrategy):
|
||||||
|
def __init__(self, t5xxl_max_length: int = 256, tokenizer_cache_dir: Optional[str] = None) -> None:
|
||||||
|
self.t5xxl_max_length = t5xxl_max_length
|
||||||
|
self.clip_l = self._load_tokenizer(CLIPTokenizer, CLIP_L_TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.t5xxl = self._load_tokenizer(T5TokenizerFast, T5_XXL_TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
|
||||||
|
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
|
||||||
|
text = [text] if isinstance(text, str) else text
|
||||||
|
|
||||||
|
l_tokens = self.clip_l(text, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
t5_tokens = self.t5xxl(text, max_length=self.t5xxl_max_length, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
|
||||||
|
t5_attn_mask = t5_tokens["attention_mask"]
|
||||||
|
l_tokens = l_tokens["input_ids"]
|
||||||
|
t5_tokens = t5_tokens["input_ids"]
|
||||||
|
|
||||||
|
return [l_tokens, t5_tokens, t5_attn_mask]
|
||||||
|
|
||||||
|
|
||||||
|
class FluxTextEncodingStrategy(TextEncodingStrategy):
|
||||||
|
def __init__(self, apply_t5_attn_mask: Optional[bool] = None) -> None:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
apply_t5_attn_mask: Default value for apply_t5_attn_mask.
|
||||||
|
"""
|
||||||
|
self.apply_t5_attn_mask = apply_t5_attn_mask
|
||||||
|
|
||||||
|
def encode_tokens(
|
||||||
|
self,
|
||||||
|
tokenize_strategy: TokenizeStrategy,
|
||||||
|
models: List[Any],
|
||||||
|
tokens: List[torch.Tensor],
|
||||||
|
apply_t5_attn_mask: Optional[bool] = None,
|
||||||
|
) -> List[torch.Tensor]:
|
||||||
|
# supports single model inference
|
||||||
|
|
||||||
|
if apply_t5_attn_mask is None:
|
||||||
|
apply_t5_attn_mask = self.apply_t5_attn_mask
|
||||||
|
|
||||||
|
clip_l, t5xxl = models
|
||||||
|
l_tokens, t5_tokens = tokens[:2]
|
||||||
|
t5_attn_mask = tokens[2] if len(tokens) > 2 else None
|
||||||
|
|
||||||
|
if clip_l is not None and l_tokens is not None:
|
||||||
|
l_pooled = clip_l(l_tokens.to(clip_l.device))["pooler_output"]
|
||||||
|
else:
|
||||||
|
l_pooled = None
|
||||||
|
|
||||||
|
if t5xxl is not None and t5_tokens is not None:
|
||||||
|
# t5_out is [b, max length, 4096]
|
||||||
|
t5_out, _ = t5xxl(t5_tokens.to(t5xxl.device), return_dict=False, output_hidden_states=True)
|
||||||
|
if apply_t5_attn_mask:
|
||||||
|
t5_out = t5_out * t5_attn_mask.to(t5_out.device).unsqueeze(-1)
|
||||||
|
txt_ids = torch.zeros(t5_out.shape[0], t5_out.shape[1], 3, device=t5_out.device)
|
||||||
|
else:
|
||||||
|
t5_out = None
|
||||||
|
txt_ids = None
|
||||||
|
|
||||||
|
return [l_pooled, t5_out, txt_ids]
|
||||||
|
|
||||||
|
|
||||||
|
class FluxTextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||||
|
FLUX_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX = "_flux_te.npz"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
cache_to_disk: bool,
|
||||||
|
batch_size: int,
|
||||||
|
skip_disk_cache_validity_check: bool,
|
||||||
|
is_partial: bool = False,
|
||||||
|
apply_t5_attn_mask: bool = False,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial)
|
||||||
|
self.apply_t5_attn_mask = apply_t5_attn_mask
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
def is_disk_cached_outputs_expected(self, npz_path: str):
|
||||||
|
if not self.cache_to_disk:
|
||||||
|
return False
|
||||||
|
if not os.path.exists(npz_path):
|
||||||
|
return False
|
||||||
|
if self.skip_disk_cache_validity_check:
|
||||||
|
return True
|
||||||
|
|
||||||
|
try:
|
||||||
|
npz = np.load(npz_path)
|
||||||
|
if "l_pooled" not in npz:
|
||||||
|
return False
|
||||||
|
if "t5_out" not in npz:
|
||||||
|
return False
|
||||||
|
if "txt_ids" not in npz:
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error loading file: {npz_path}")
|
||||||
|
raise e
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
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)
|
||||||
|
l_pooled = data["l_pooled"]
|
||||||
|
t5_out = data["t5_out"]
|
||||||
|
txt_ids = data["txt_ids"]
|
||||||
|
|
||||||
|
if self.apply_t5_attn_mask:
|
||||||
|
t5_attn_mask = data["t5_attn_mask"]
|
||||||
|
t5_out = self.mask_t5_attn(t5_out, t5_attn_mask)
|
||||||
|
|
||||||
|
return [l_pooled, t5_out, txt_ids]
|
||||||
|
|
||||||
|
def cache_batch_outputs(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
|
||||||
|
):
|
||||||
|
flux_text_encoding_strategy: FluxTextEncodingStrategy = text_encoding_strategy
|
||||||
|
captions = [info.caption for info in infos]
|
||||||
|
|
||||||
|
tokens_and_masks = tokenize_strategy.tokenize(captions)
|
||||||
|
with torch.no_grad():
|
||||||
|
# attn_mask is not applied when caching to disk: it is applied when loading from disk
|
||||||
|
l_pooled, t5_out, txt_ids = flux_text_encoding_strategy.encode_tokens(
|
||||||
|
tokenize_strategy, models, tokens_and_masks, not self.cache_to_disk
|
||||||
|
)
|
||||||
|
|
||||||
|
if l_pooled.dtype == torch.bfloat16:
|
||||||
|
l_pooled = l_pooled.float()
|
||||||
|
if t5_out.dtype == torch.bfloat16:
|
||||||
|
t5_out = t5_out.float()
|
||||||
|
if txt_ids.dtype == torch.bfloat16:
|
||||||
|
txt_ids = txt_ids.float()
|
||||||
|
|
||||||
|
l_pooled = l_pooled.cpu().numpy()
|
||||||
|
t5_out = t5_out.cpu().numpy()
|
||||||
|
txt_ids = txt_ids.cpu().numpy()
|
||||||
|
|
||||||
|
for i, info in enumerate(infos):
|
||||||
|
l_pooled_i = l_pooled[i]
|
||||||
|
t5_out_i = t5_out[i]
|
||||||
|
txt_ids_i = txt_ids[i]
|
||||||
|
|
||||||
|
if self.cache_to_disk:
|
||||||
|
t5_attn_mask = tokens_and_masks[2]
|
||||||
|
t5_attn_mask_i = t5_attn_mask[i].cpu().numpy()
|
||||||
|
np.savez(
|
||||||
|
info.text_encoder_outputs_npz,
|
||||||
|
l_pooled=l_pooled_i,
|
||||||
|
t5_out=t5_out_i,
|
||||||
|
txt_ids=txt_ids_i,
|
||||||
|
t5_attn_mask=t5_attn_mask_i,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
info.text_encoder_outputs = (l_pooled_i, t5_out_i, txt_ids_i)
|
||||||
|
|
||||||
|
|
||||||
|
class FluxLatentsCachingStrategy(LatentsCachingStrategy):
|
||||||
|
FLUX_LATENTS_NPZ_SUFFIX = "_flux.npz"
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||||
|
return (
|
||||||
|
os.path.splitext(absolute_path)[0]
|
||||||
|
+ f"_{image_size[0]:04d}x{image_size[1]:04d}"
|
||||||
|
+ FluxLatentsCachingStrategy.FLUX_LATENTS_NPZ_SUFFIX
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
# TODO remove circular dependency for ImageInfo
|
||||||
|
def cache_batch_latents(self, vae, image_infos: List, flip_aug: bool, alpha_mask: bool, random_crop: bool):
|
||||||
|
encode_by_vae = lambda img_tensor: vae.encode(img_tensor).to("cpu")
|
||||||
|
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)
|
||||||
|
|
||||||
|
if not train_util.HIGH_VRAM:
|
||||||
|
train_util.clean_memory_on_device(vae.device)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# test code for FluxTokenizeStrategy
|
||||||
|
# tokenizer = sd3_models.SD3Tokenizer()
|
||||||
|
strategy = FluxTokenizeStrategy(256)
|
||||||
|
text = "hello world"
|
||||||
|
|
||||||
|
l_tokens, g_tokens, t5_tokens = strategy.tokenize(text)
|
||||||
|
# print(l_tokens.shape)
|
||||||
|
print(l_tokens)
|
||||||
|
print(g_tokens)
|
||||||
|
print(t5_tokens)
|
||||||
|
|
||||||
|
texts = ["hello world", "the quick brown fox jumps over the lazy dog"]
|
||||||
|
l_tokens_2 = strategy.clip_l(texts, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
g_tokens_2 = strategy.clip_g(texts, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
t5_tokens_2 = strategy.t5xxl(
|
||||||
|
texts, max_length=strategy.t5xxl_max_length, padding="max_length", truncation=True, return_tensors="pt"
|
||||||
|
)
|
||||||
|
print(l_tokens_2)
|
||||||
|
print(g_tokens_2)
|
||||||
|
print(t5_tokens_2)
|
||||||
|
|
||||||
|
# compare
|
||||||
|
print(torch.allclose(l_tokens, l_tokens_2["input_ids"][0]))
|
||||||
|
print(torch.allclose(g_tokens, g_tokens_2["input_ids"][0]))
|
||||||
|
print(torch.allclose(t5_tokens, t5_tokens_2["input_ids"][0]))
|
||||||
|
|
||||||
|
text = ",".join(["hello world! this is long text"] * 50)
|
||||||
|
l_tokens, g_tokens, t5_tokens = strategy.tokenize(text)
|
||||||
|
print(l_tokens)
|
||||||
|
print(g_tokens)
|
||||||
|
print(t5_tokens)
|
||||||
|
|
||||||
|
print(f"model max length l: {strategy.clip_l.model_max_length}")
|
||||||
|
print(f"model max length g: {strategy.clip_g.model_max_length}")
|
||||||
|
print(f"model max length t5: {strategy.t5xxl.model_max_length}")
|
||||||
@@ -0,0 +1,139 @@
|
|||||||
|
import glob
|
||||||
|
import os
|
||||||
|
from typing import Any, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from transformers import CLIPTokenizer
|
||||||
|
from . import train_util
|
||||||
|
from .strategy_base import LatentsCachingStrategy, TokenizeStrategy, TextEncodingStrategy
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
TOKENIZER_ID = "openai/clip-vit-large-patch14"
|
||||||
|
V2_STABLE_DIFFUSION_ID = "stabilityai/stable-diffusion-2" # ここからtokenizerだけ使う v2とv2.1はtokenizer仕様は同じ
|
||||||
|
|
||||||
|
|
||||||
|
class SdTokenizeStrategy(TokenizeStrategy):
|
||||||
|
def __init__(self, v2: bool, max_length: Optional[int], tokenizer_cache_dir: Optional[str] = None) -> None:
|
||||||
|
"""
|
||||||
|
max_length does not include <BOS> and <EOS> (None, 75, 150, 225)
|
||||||
|
"""
|
||||||
|
logger.info(f"Using {'v2' if v2 else 'v1'} tokenizer")
|
||||||
|
if v2:
|
||||||
|
self.tokenizer = self._load_tokenizer(
|
||||||
|
CLIPTokenizer, V2_STABLE_DIFFUSION_ID, subfolder="tokenizer", tokenizer_cache_dir=tokenizer_cache_dir
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.tokenizer = self._load_tokenizer(CLIPTokenizer, TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
|
||||||
|
if max_length is None:
|
||||||
|
self.max_length = self.tokenizer.model_max_length
|
||||||
|
else:
|
||||||
|
self.max_length = max_length + 2
|
||||||
|
|
||||||
|
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
|
||||||
|
text = [text] if isinstance(text, str) else text
|
||||||
|
return [torch.stack([self._get_input_ids(self.tokenizer, t, self.max_length) for t in text], dim=0)]
|
||||||
|
|
||||||
|
|
||||||
|
class SdTextEncodingStrategy(TextEncodingStrategy):
|
||||||
|
def __init__(self, clip_skip: Optional[int] = None) -> None:
|
||||||
|
self.clip_skip = clip_skip
|
||||||
|
|
||||||
|
def encode_tokens(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], tokens: List[torch.Tensor]
|
||||||
|
) -> List[torch.Tensor]:
|
||||||
|
text_encoder = models[0]
|
||||||
|
tokens = tokens[0]
|
||||||
|
sd_tokenize_strategy = tokenize_strategy # type: SdTokenizeStrategy
|
||||||
|
|
||||||
|
# tokens: b,n,77
|
||||||
|
b_size = tokens.size()[0]
|
||||||
|
max_token_length = tokens.size()[1] * tokens.size()[2]
|
||||||
|
model_max_length = sd_tokenize_strategy.tokenizer.model_max_length
|
||||||
|
tokens = tokens.reshape((-1, model_max_length)) # batch_size*3, 77
|
||||||
|
|
||||||
|
if self.clip_skip is None:
|
||||||
|
encoder_hidden_states = text_encoder(tokens)[0]
|
||||||
|
else:
|
||||||
|
enc_out = text_encoder(tokens, output_hidden_states=True, return_dict=True)
|
||||||
|
encoder_hidden_states = enc_out["hidden_states"][-self.clip_skip]
|
||||||
|
encoder_hidden_states = text_encoder.text_model.final_layer_norm(encoder_hidden_states)
|
||||||
|
|
||||||
|
# bs*3, 77, 768 or 1024
|
||||||
|
encoder_hidden_states = encoder_hidden_states.reshape((b_size, -1, encoder_hidden_states.shape[-1]))
|
||||||
|
|
||||||
|
if max_token_length != model_max_length:
|
||||||
|
v1 = sd_tokenize_strategy.tokenizer.pad_token_id == sd_tokenize_strategy.tokenizer.eos_token_id
|
||||||
|
if not v1:
|
||||||
|
# v2: <BOS>...<EOS> <PAD> ... の三連を <BOS>...<EOS> <PAD> ... へ戻す 正直この実装でいいのかわからん
|
||||||
|
states_list = [encoder_hidden_states[:, 0].unsqueeze(1)] # <BOS>
|
||||||
|
for i in range(1, max_token_length, model_max_length):
|
||||||
|
chunk = encoder_hidden_states[:, i : i + model_max_length - 2] # <BOS> の後から 最後の前まで
|
||||||
|
if i > 0:
|
||||||
|
for j in range(len(chunk)):
|
||||||
|
if tokens[j, 1] == sd_tokenize_strategy.tokenizer.eos_token:
|
||||||
|
# 空、つまり <BOS> <EOS> <PAD> ...のパターン
|
||||||
|
chunk[j, 0] = chunk[j, 1] # 次の <PAD> の値をコピーする
|
||||||
|
states_list.append(chunk) # <BOS> の後から <EOS> の前まで
|
||||||
|
states_list.append(encoder_hidden_states[:, -1].unsqueeze(1)) # <EOS> か <PAD> のどちらか
|
||||||
|
encoder_hidden_states = torch.cat(states_list, dim=1)
|
||||||
|
else:
|
||||||
|
# v1: <BOS>...<EOS> の三連を <BOS>...<EOS> へ戻す
|
||||||
|
states_list = [encoder_hidden_states[:, 0].unsqueeze(1)] # <BOS>
|
||||||
|
for i in range(1, max_token_length, model_max_length):
|
||||||
|
states_list.append(encoder_hidden_states[:, i : i + model_max_length - 2]) # <BOS> の後から <EOS> の前まで
|
||||||
|
states_list.append(encoder_hidden_states[:, -1].unsqueeze(1)) # <EOS>
|
||||||
|
encoder_hidden_states = torch.cat(states_list, dim=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.
|
||||||
|
# and we keep the old npz for the backward compatibility.
|
||||||
|
|
||||||
|
SD_OLD_LATENTS_NPZ_SUFFIX = ".npz"
|
||||||
|
SD_LATENTS_NPZ_SUFFIX = "_sd.npz"
|
||||||
|
SDXL_LATENTS_NPZ_SUFFIX = "_sdxl.npz"
|
||||||
|
|
||||||
|
def __init__(self, sd: bool, 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)
|
||||||
|
self.sd = sd
|
||||||
|
self.suffix = (
|
||||||
|
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)
|
||||||
|
|
||||||
|
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||||
|
# support old .npz
|
||||||
|
old_npz_file = os.path.splitext(absolute_path)[0] + SdSdxlLatentsCachingStrategy.SD_OLD_LATENTS_NPZ_SUFFIX
|
||||||
|
if os.path.exists(old_npz_file):
|
||||||
|
return old_npz_file
|
||||||
|
return os.path.splitext(absolute_path)[0] + f"_{image_size[0]:04d}x{image_size[1]:04d}" + self.suffix
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
# TODO remove circular dependency for ImageInfo
|
||||||
|
def cache_batch_latents(self, vae, image_infos: List, flip_aug: bool, alpha_mask: bool, random_crop: bool):
|
||||||
|
encode_by_vae = lambda img_tensor: vae.encode(img_tensor).latent_dist.sample()
|
||||||
|
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)
|
||||||
|
|
||||||
|
if not train_util.HIGH_VRAM:
|
||||||
|
train_util.clean_memory_on_device(vae.device)
|
||||||
@@ -0,0 +1,289 @@
|
|||||||
|
import os
|
||||||
|
import glob
|
||||||
|
from typing import Any, List, Optional, Tuple, Union
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
from transformers import CLIPTokenizer, T5TokenizerFast
|
||||||
|
|
||||||
|
from library import sd3_utils, train_util
|
||||||
|
from library import sd3_models
|
||||||
|
from library.strategy_base import LatentsCachingStrategy, TextEncodingStrategy, TokenizeStrategy, TextEncoderOutputsCachingStrategy
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
CLIP_L_TOKENIZER_ID = "openai/clip-vit-large-patch14"
|
||||||
|
CLIP_G_TOKENIZER_ID = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k"
|
||||||
|
T5_XXL_TOKENIZER_ID = "google/t5-v1_1-xxl"
|
||||||
|
|
||||||
|
|
||||||
|
class Sd3TokenizeStrategy(TokenizeStrategy):
|
||||||
|
def __init__(self, t5xxl_max_length: int = 256, tokenizer_cache_dir: Optional[str] = None) -> None:
|
||||||
|
self.t5xxl_max_length = t5xxl_max_length
|
||||||
|
self.clip_l = self._load_tokenizer(CLIPTokenizer, CLIP_L_TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.clip_g = self._load_tokenizer(CLIPTokenizer, CLIP_G_TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.t5xxl = self._load_tokenizer(T5TokenizerFast, T5_XXL_TOKENIZER_ID, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.clip_g.pad_token_id = 0 # use 0 as pad token for clip_g
|
||||||
|
|
||||||
|
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
|
||||||
|
text = [text] if isinstance(text, str) else text
|
||||||
|
|
||||||
|
l_tokens = self.clip_l(text, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
g_tokens = self.clip_g(text, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
t5_tokens = self.t5xxl(text, max_length=self.t5xxl_max_length, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
|
||||||
|
l_attn_mask = l_tokens["attention_mask"]
|
||||||
|
g_attn_mask = g_tokens["attention_mask"]
|
||||||
|
t5_attn_mask = t5_tokens["attention_mask"]
|
||||||
|
l_tokens = l_tokens["input_ids"]
|
||||||
|
g_tokens = g_tokens["input_ids"]
|
||||||
|
t5_tokens = t5_tokens["input_ids"]
|
||||||
|
|
||||||
|
return [l_tokens, g_tokens, t5_tokens, l_attn_mask, g_attn_mask, t5_attn_mask]
|
||||||
|
|
||||||
|
|
||||||
|
class Sd3TextEncodingStrategy(TextEncodingStrategy):
|
||||||
|
def __init__(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
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,
|
||||||
|
) -> List[torch.Tensor]:
|
||||||
|
"""
|
||||||
|
returned embeddings are not masked
|
||||||
|
"""
|
||||||
|
clip_l, clip_g, t5xxl = models
|
||||||
|
|
||||||
|
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:
|
||||||
|
assert g_tokens is None, "g_tokens must be None if l_tokens is None"
|
||||||
|
lg_out = 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)
|
||||||
|
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:
|
||||||
|
t5_out = None
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
def concat_encodings(
|
||||||
|
self, lg_out: torch.Tensor, t5_out: Optional[torch.Tensor], lg_pooled: torch.Tensor
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
lg_out = torch.nn.functional.pad(lg_out, (0, 4096 - lg_out.shape[-1]))
|
||||||
|
if t5_out is None:
|
||||||
|
t5_out = torch.zeros((lg_out.shape[0], 77, 4096), device=lg_out.device, dtype=lg_out.dtype)
|
||||||
|
return torch.cat([lg_out, t5_out], dim=-2), lg_pooled
|
||||||
|
|
||||||
|
|
||||||
|
class Sd3TextEncoderOutputsCachingStrategy(TextEncoderOutputsCachingStrategy):
|
||||||
|
SD3_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX = "_sd3_te.npz"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
cache_to_disk: bool,
|
||||||
|
batch_size: int,
|
||||||
|
skip_disk_cache_validity_check: bool,
|
||||||
|
is_partial: bool = False,
|
||||||
|
apply_lg_attn_mask: bool = False,
|
||||||
|
apply_t5_attn_mask: bool = False,
|
||||||
|
) -> None:
|
||||||
|
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial)
|
||||||
|
self.apply_lg_attn_mask = apply_lg_attn_mask
|
||||||
|
self.apply_t5_attn_mask = apply_t5_attn_mask
|
||||||
|
|
||||||
|
def get_outputs_npz_path(self, image_abs_path: str) -> str:
|
||||||
|
return os.path.splitext(image_abs_path)[0] + Sd3TextEncoderOutputsCachingStrategy.SD3_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX
|
||||||
|
|
||||||
|
def is_disk_cached_outputs_expected(self, npz_path: str):
|
||||||
|
if not self.cache_to_disk:
|
||||||
|
return False
|
||||||
|
if not os.path.exists(npz_path):
|
||||||
|
return False
|
||||||
|
if self.skip_disk_cache_validity_check:
|
||||||
|
return True
|
||||||
|
|
||||||
|
try:
|
||||||
|
npz = np.load(npz_path)
|
||||||
|
if "lg_out" not in npz:
|
||||||
|
return False
|
||||||
|
if "lg_pooled" not in npz:
|
||||||
|
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
|
||||||
|
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
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
def cache_batch_outputs(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
|
||||||
|
):
|
||||||
|
sd3_text_encoding_strategy: Sd3TextEncodingStrategy = text_encoding_strategy
|
||||||
|
captions = [info.caption for info in infos]
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
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:
|
||||||
|
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()
|
||||||
|
|
||||||
|
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
|
||||||
|
lg_pooled_i = lg_pooled[i]
|
||||||
|
|
||||||
|
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_attn_mask=t5_attn_mask_i,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
info.text_encoder_outputs = (lg_out_i, t5_out_i, lg_pooled_i)
|
||||||
|
|
||||||
|
|
||||||
|
class Sd3LatentsCachingStrategy(LatentsCachingStrategy):
|
||||||
|
SD3_LATENTS_NPZ_SUFFIX = "_sd3.npz"
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
def get_latents_npz_path(self, absolute_path: str, image_size: Tuple[int, int]) -> str:
|
||||||
|
return (
|
||||||
|
os.path.splitext(absolute_path)[0]
|
||||||
|
+ f"_{image_size[0]:04d}x{image_size[1]:04d}"
|
||||||
|
+ Sd3LatentsCachingStrategy.SD3_LATENTS_NPZ_SUFFIX
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
# TODO remove circular dependency for ImageInfo
|
||||||
|
def cache_batch_latents(self, vae, image_infos: List, flip_aug: bool, alpha_mask: bool, random_crop: bool):
|
||||||
|
encode_by_vae = lambda img_tensor: vae.encode(img_tensor).to("cpu")
|
||||||
|
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)
|
||||||
|
|
||||||
|
if not train_util.HIGH_VRAM:
|
||||||
|
train_util.clean_memory_on_device(vae.device)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# test code for Sd3TokenizeStrategy
|
||||||
|
# tokenizer = sd3_models.SD3Tokenizer()
|
||||||
|
strategy = Sd3TokenizeStrategy(256)
|
||||||
|
text = "hello world"
|
||||||
|
|
||||||
|
l_tokens, g_tokens, t5_tokens = strategy.tokenize(text)
|
||||||
|
# print(l_tokens.shape)
|
||||||
|
print(l_tokens)
|
||||||
|
print(g_tokens)
|
||||||
|
print(t5_tokens)
|
||||||
|
|
||||||
|
texts = ["hello world", "the quick brown fox jumps over the lazy dog"]
|
||||||
|
l_tokens_2 = strategy.clip_l(texts, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
g_tokens_2 = strategy.clip_g(texts, max_length=77, padding="max_length", truncation=True, return_tensors="pt")
|
||||||
|
t5_tokens_2 = strategy.t5xxl(
|
||||||
|
texts, max_length=strategy.t5xxl_max_length, padding="max_length", truncation=True, return_tensors="pt"
|
||||||
|
)
|
||||||
|
print(l_tokens_2)
|
||||||
|
print(g_tokens_2)
|
||||||
|
print(t5_tokens_2)
|
||||||
|
|
||||||
|
# compare
|
||||||
|
print(torch.allclose(l_tokens, l_tokens_2["input_ids"][0]))
|
||||||
|
print(torch.allclose(g_tokens, g_tokens_2["input_ids"][0]))
|
||||||
|
print(torch.allclose(t5_tokens, t5_tokens_2["input_ids"][0]))
|
||||||
|
|
||||||
|
text = ",".join(["hello world! this is long text"] * 50)
|
||||||
|
l_tokens, g_tokens, t5_tokens = strategy.tokenize(text)
|
||||||
|
print(l_tokens)
|
||||||
|
print(g_tokens)
|
||||||
|
print(t5_tokens)
|
||||||
|
|
||||||
|
print(f"model max length l: {strategy.clip_l.model_max_length}")
|
||||||
|
print(f"model max length g: {strategy.clip_g.model_max_length}")
|
||||||
|
print(f"model max length t5: {strategy.t5xxl.model_max_length}")
|
||||||
@@ -0,0 +1,247 @@
|
|||||||
|
import os
|
||||||
|
from typing import Any, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from transformers import CLIPTokenizer, CLIPTextModel, CLIPTextModelWithProjection
|
||||||
|
from library.strategy_base import TokenizeStrategy, TextEncodingStrategy, TextEncoderOutputsCachingStrategy
|
||||||
|
|
||||||
|
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
TOKENIZER1_PATH = "openai/clip-vit-large-patch14"
|
||||||
|
TOKENIZER2_PATH = "laion/CLIP-ViT-bigG-14-laion2B-39B-b160k"
|
||||||
|
|
||||||
|
|
||||||
|
class SdxlTokenizeStrategy(TokenizeStrategy):
|
||||||
|
def __init__(self, max_length: Optional[int], tokenizer_cache_dir: Optional[str] = None) -> None:
|
||||||
|
self.tokenizer1 = self._load_tokenizer(CLIPTokenizer, TOKENIZER1_PATH, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.tokenizer2 = self._load_tokenizer(CLIPTokenizer, TOKENIZER2_PATH, tokenizer_cache_dir=tokenizer_cache_dir)
|
||||||
|
self.tokenizer2.pad_token_id = 0 # use 0 as pad token for tokenizer2
|
||||||
|
|
||||||
|
if max_length is None:
|
||||||
|
self.max_length = self.tokenizer1.model_max_length
|
||||||
|
else:
|
||||||
|
self.max_length = max_length + 2
|
||||||
|
|
||||||
|
def tokenize(self, text: Union[str, List[str]]) -> List[torch.Tensor]:
|
||||||
|
text = [text] if isinstance(text, str) else text
|
||||||
|
return (
|
||||||
|
torch.stack([self._get_input_ids(self.tokenizer1, t, self.max_length) for t in text], dim=0),
|
||||||
|
torch.stack([self._get_input_ids(self.tokenizer2, t, self.max_length) for t in text], dim=0),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class SdxlTextEncodingStrategy(TextEncodingStrategy):
|
||||||
|
def __init__(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _pool_workaround(
|
||||||
|
self, text_encoder: CLIPTextModelWithProjection, last_hidden_state: torch.Tensor, input_ids: torch.Tensor, eos_token_id: int
|
||||||
|
):
|
||||||
|
r"""
|
||||||
|
workaround for CLIP's pooling bug: it returns the hidden states for the max token id as the pooled output
|
||||||
|
instead of the hidden states for the EOS token
|
||||||
|
If we use Textual Inversion, we need to use the hidden states for the EOS token as the pooled output
|
||||||
|
|
||||||
|
Original code from CLIP's pooling function:
|
||||||
|
|
||||||
|
\# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||||
|
\# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||||
|
\# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
|
||||||
|
pooled_output = last_hidden_state[
|
||||||
|
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||||
|
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
|
||||||
|
]
|
||||||
|
"""
|
||||||
|
|
||||||
|
# input_ids: b*n,77
|
||||||
|
# find index for EOS token
|
||||||
|
|
||||||
|
# Following code is not working if one of the input_ids has multiple EOS tokens (very odd case)
|
||||||
|
# eos_token_index = torch.where(input_ids == eos_token_id)[1]
|
||||||
|
# eos_token_index = eos_token_index.to(device=last_hidden_state.device)
|
||||||
|
|
||||||
|
# Create a mask where the EOS tokens are
|
||||||
|
eos_token_mask = (input_ids == eos_token_id).int()
|
||||||
|
|
||||||
|
# Use argmax to find the last index of the EOS token for each element in the batch
|
||||||
|
eos_token_index = torch.argmax(eos_token_mask, dim=1) # this will be 0 if there is no EOS token, it's fine
|
||||||
|
eos_token_index = eos_token_index.to(device=last_hidden_state.device)
|
||||||
|
|
||||||
|
# get hidden states for EOS token
|
||||||
|
pooled_output = last_hidden_state[
|
||||||
|
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device), eos_token_index
|
||||||
|
]
|
||||||
|
|
||||||
|
# apply projection: projection may be of different dtype than last_hidden_state
|
||||||
|
pooled_output = text_encoder.text_projection(pooled_output.to(text_encoder.text_projection.weight.dtype))
|
||||||
|
pooled_output = pooled_output.to(last_hidden_state.dtype)
|
||||||
|
|
||||||
|
return pooled_output
|
||||||
|
|
||||||
|
def _get_hidden_states_sdxl(
|
||||||
|
self,
|
||||||
|
input_ids1: torch.Tensor,
|
||||||
|
input_ids2: torch.Tensor,
|
||||||
|
tokenizer1: CLIPTokenizer,
|
||||||
|
tokenizer2: CLIPTokenizer,
|
||||||
|
text_encoder1: Union[CLIPTextModel, torch.nn.Module],
|
||||||
|
text_encoder2: Union[CLIPTextModelWithProjection, torch.nn.Module],
|
||||||
|
unwrapped_text_encoder2: Optional[CLIPTextModelWithProjection] = None,
|
||||||
|
):
|
||||||
|
# input_ids: b,n,77 -> b*n, 77
|
||||||
|
b_size = input_ids1.size()[0]
|
||||||
|
max_token_length = input_ids1.size()[1] * input_ids1.size()[2]
|
||||||
|
input_ids1 = input_ids1.reshape((-1, tokenizer1.model_max_length)) # batch_size*n, 77
|
||||||
|
input_ids2 = input_ids2.reshape((-1, tokenizer2.model_max_length)) # batch_size*n, 77
|
||||||
|
input_ids1 = input_ids1.to(text_encoder1.device)
|
||||||
|
input_ids2 = input_ids2.to(text_encoder2.device)
|
||||||
|
|
||||||
|
# text_encoder1
|
||||||
|
enc_out = text_encoder1(input_ids1, output_hidden_states=True, return_dict=True)
|
||||||
|
hidden_states1 = enc_out["hidden_states"][11]
|
||||||
|
|
||||||
|
# text_encoder2
|
||||||
|
enc_out = text_encoder2(input_ids2, output_hidden_states=True, return_dict=True)
|
||||||
|
hidden_states2 = enc_out["hidden_states"][-2] # penuultimate layer
|
||||||
|
|
||||||
|
# pool2 = enc_out["text_embeds"]
|
||||||
|
unwrapped_text_encoder2 = unwrapped_text_encoder2 or text_encoder2
|
||||||
|
pool2 = self._pool_workaround(unwrapped_text_encoder2, enc_out["last_hidden_state"], input_ids2, tokenizer2.eos_token_id)
|
||||||
|
|
||||||
|
# b*n, 77, 768 or 1280 -> b, n*77, 768 or 1280
|
||||||
|
n_size = 1 if max_token_length is None else max_token_length // 75
|
||||||
|
hidden_states1 = hidden_states1.reshape((b_size, -1, hidden_states1.shape[-1]))
|
||||||
|
hidden_states2 = hidden_states2.reshape((b_size, -1, hidden_states2.shape[-1]))
|
||||||
|
|
||||||
|
if max_token_length is not None:
|
||||||
|
# bs*3, 77, 768 or 1024
|
||||||
|
# encoder1: <BOS>...<EOS> の三連を <BOS>...<EOS> へ戻す
|
||||||
|
states_list = [hidden_states1[:, 0].unsqueeze(1)] # <BOS>
|
||||||
|
for i in range(1, max_token_length, tokenizer1.model_max_length):
|
||||||
|
states_list.append(hidden_states1[:, i : i + tokenizer1.model_max_length - 2]) # <BOS> の後から <EOS> の前まで
|
||||||
|
states_list.append(hidden_states1[:, -1].unsqueeze(1)) # <EOS>
|
||||||
|
hidden_states1 = torch.cat(states_list, dim=1)
|
||||||
|
|
||||||
|
# v2: <BOS>...<EOS> <PAD> ... の三連を <BOS>...<EOS> <PAD> ... へ戻す 正直この実装でいいのかわからん
|
||||||
|
states_list = [hidden_states2[:, 0].unsqueeze(1)] # <BOS>
|
||||||
|
for i in range(1, max_token_length, tokenizer2.model_max_length):
|
||||||
|
chunk = hidden_states2[:, i : i + tokenizer2.model_max_length - 2] # <BOS> の後から 最後の前まで
|
||||||
|
# this causes an error:
|
||||||
|
# RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation
|
||||||
|
# if i > 1:
|
||||||
|
# for j in range(len(chunk)): # batch_size
|
||||||
|
# if input_ids2[n_index + j * n_size, 1] == tokenizer2.eos_token_id: # 空、つまり <BOS> <EOS> <PAD> ...のパターン
|
||||||
|
# chunk[j, 0] = chunk[j, 1] # 次の <PAD> の値をコピーする
|
||||||
|
states_list.append(chunk) # <BOS> の後から <EOS> の前まで
|
||||||
|
states_list.append(hidden_states2[:, -1].unsqueeze(1)) # <EOS> か <PAD> のどちらか
|
||||||
|
hidden_states2 = torch.cat(states_list, dim=1)
|
||||||
|
|
||||||
|
# pool はnの最初のものを使う
|
||||||
|
pool2 = pool2[::n_size]
|
||||||
|
|
||||||
|
return hidden_states1, hidden_states2, pool2
|
||||||
|
|
||||||
|
def encode_tokens(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], tokens: List[torch.Tensor]
|
||||||
|
) -> List[torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
tokenize_strategy: TokenizeStrategy
|
||||||
|
models: List of models, [text_encoder1, text_encoder2, unwrapped text_encoder2 (optional)]
|
||||||
|
tokens: List of tokens, for text_encoder1 and text_encoder2
|
||||||
|
"""
|
||||||
|
if len(models) == 2:
|
||||||
|
text_encoder1, text_encoder2 = models
|
||||||
|
unwrapped_text_encoder2 = None
|
||||||
|
else:
|
||||||
|
text_encoder1, text_encoder2, unwrapped_text_encoder2 = models
|
||||||
|
tokens1, tokens2 = tokens
|
||||||
|
sdxl_tokenize_strategy = tokenize_strategy # type: SdxlTokenizeStrategy
|
||||||
|
tokenizer1, tokenizer2 = sdxl_tokenize_strategy.tokenizer1, sdxl_tokenize_strategy.tokenizer2
|
||||||
|
|
||||||
|
hidden_states1, hidden_states2, pool2 = self._get_hidden_states_sdxl(
|
||||||
|
tokens1, tokens2, tokenizer1, tokenizer2, text_encoder1, text_encoder2, unwrapped_text_encoder2
|
||||||
|
)
|
||||||
|
return [hidden_states1, hidden_states2, pool2]
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
) -> None:
|
||||||
|
super().__init__(cache_to_disk, batch_size, skip_disk_cache_validity_check, is_partial)
|
||||||
|
|
||||||
|
def get_outputs_npz_path(self, image_abs_path: str) -> str:
|
||||||
|
return os.path.splitext(image_abs_path)[0] + SdxlTextEncoderOutputsCachingStrategy.SDXL_TEXT_ENCODER_OUTPUTS_NPZ_SUFFIX
|
||||||
|
|
||||||
|
def is_disk_cached_outputs_expected(self, npz_path: str):
|
||||||
|
if not self.cache_to_disk:
|
||||||
|
return False
|
||||||
|
if not os.path.exists(npz_path):
|
||||||
|
return False
|
||||||
|
if self.skip_disk_cache_validity_check:
|
||||||
|
return True
|
||||||
|
|
||||||
|
try:
|
||||||
|
npz = np.load(npz_path)
|
||||||
|
if "hidden_state1" not in npz or "hidden_state2" not in npz or "pool2" not in npz:
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error loading file: {npz_path}")
|
||||||
|
raise e
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def load_outputs_npz(self, npz_path: str) -> List[np.ndarray]:
|
||||||
|
data = np.load(npz_path)
|
||||||
|
hidden_state1 = data["hidden_state1"]
|
||||||
|
hidden_state2 = data["hidden_state2"]
|
||||||
|
pool2 = data["pool2"]
|
||||||
|
return [hidden_state1, hidden_state2, pool2]
|
||||||
|
|
||||||
|
def cache_batch_outputs(
|
||||||
|
self, tokenize_strategy: TokenizeStrategy, models: List[Any], text_encoding_strategy: TextEncodingStrategy, infos: List
|
||||||
|
):
|
||||||
|
sdxl_text_encoding_strategy = text_encoding_strategy # type: SdxlTextEncodingStrategy
|
||||||
|
captions = [info.caption for info in infos]
|
||||||
|
|
||||||
|
tokens1, tokens2 = tokenize_strategy.tokenize(captions)
|
||||||
|
with torch.no_grad():
|
||||||
|
hidden_state1, hidden_state2, pool2 = sdxl_text_encoding_strategy.encode_tokens(
|
||||||
|
tokenize_strategy, models, [tokens1, tokens2]
|
||||||
|
)
|
||||||
|
if hidden_state1.dtype == torch.bfloat16:
|
||||||
|
hidden_state1 = hidden_state1.float()
|
||||||
|
if hidden_state2.dtype == torch.bfloat16:
|
||||||
|
hidden_state2 = hidden_state2.float()
|
||||||
|
if pool2.dtype == torch.bfloat16:
|
||||||
|
pool2 = pool2.float()
|
||||||
|
|
||||||
|
hidden_state1 = hidden_state1.cpu().numpy()
|
||||||
|
hidden_state2 = hidden_state2.cpu().numpy()
|
||||||
|
pool2 = pool2.cpu().numpy()
|
||||||
|
|
||||||
|
for i, info in enumerate(infos):
|
||||||
|
hidden_state1_i = hidden_state1[i]
|
||||||
|
hidden_state2_i = hidden_state2[i]
|
||||||
|
pool2_i = pool2[i]
|
||||||
|
|
||||||
|
if self.cache_to_disk:
|
||||||
|
np.savez(
|
||||||
|
info.text_encoder_outputs_npz,
|
||||||
|
hidden_state1=hidden_state1_i,
|
||||||
|
hidden_state2=hidden_state2_i,
|
||||||
|
pool2=pool2_i,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
info.text_encoder_outputs = [hidden_state1_i, hidden_state2_i, pool2_i]
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,266 @@
|
|||||||
|
import logging
|
||||||
|
import sys
|
||||||
|
import threading
|
||||||
|
import torch
|
||||||
|
from torchvision import transforms
|
||||||
|
from typing import *
|
||||||
|
from diffusers import EulerAncestralDiscreteScheduler
|
||||||
|
import diffusers.schedulers.scheduling_euler_ancestral_discrete
|
||||||
|
from diffusers.schedulers.scheduling_euler_ancestral_discrete import EulerAncestralDiscreteSchedulerOutput
|
||||||
|
|
||||||
|
|
||||||
|
def fire_in_thread(f, *args, **kwargs):
|
||||||
|
threading.Thread(target=f, args=args, kwargs=kwargs).start()
|
||||||
|
|
||||||
|
|
||||||
|
def add_logging_arguments(parser):
|
||||||
|
parser.add_argument(
|
||||||
|
"--console_log_level",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"],
|
||||||
|
help="Set the logging level, default is INFO / ログレベルを設定する。デフォルトはINFO",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--console_log_file",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Log to a file instead of stderr / 標準エラー出力ではなくファイルにログを出力する",
|
||||||
|
)
|
||||||
|
parser.add_argument("--console_log_simple", action="store_true", help="Simple log output / シンプルなログ出力")
|
||||||
|
|
||||||
|
|
||||||
|
def setup_logging(args=None, log_level=None, reset=False):
|
||||||
|
if logging.root.handlers:
|
||||||
|
if reset:
|
||||||
|
# remove all handlers
|
||||||
|
for handler in logging.root.handlers[:]:
|
||||||
|
logging.root.removeHandler(handler)
|
||||||
|
else:
|
||||||
|
return
|
||||||
|
|
||||||
|
# log_level can be set by the caller or by the args, the caller has priority. If not set, use INFO
|
||||||
|
if log_level is None and args is not None:
|
||||||
|
log_level = args.console_log_level
|
||||||
|
if log_level is None:
|
||||||
|
log_level = "INFO"
|
||||||
|
log_level = getattr(logging, log_level)
|
||||||
|
|
||||||
|
msg_init = None
|
||||||
|
if args is not None and args.console_log_file:
|
||||||
|
handler = logging.FileHandler(args.console_log_file, mode="w")
|
||||||
|
else:
|
||||||
|
handler = None
|
||||||
|
if not args or not args.console_log_simple:
|
||||||
|
try:
|
||||||
|
from rich.logging import RichHandler
|
||||||
|
from rich.console import Console
|
||||||
|
from rich.logging import RichHandler
|
||||||
|
|
||||||
|
handler = RichHandler(console=Console(stderr=True))
|
||||||
|
except ImportError:
|
||||||
|
# print("rich is not installed, using basic logging")
|
||||||
|
msg_init = "rich is not installed, using basic logging"
|
||||||
|
|
||||||
|
if handler is None:
|
||||||
|
handler = logging.StreamHandler(sys.stdout) # same as print
|
||||||
|
handler.propagate = False
|
||||||
|
|
||||||
|
formatter = logging.Formatter(
|
||||||
|
fmt="%(message)s",
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
handler.setFormatter(formatter)
|
||||||
|
logging.root.setLevel(log_level)
|
||||||
|
logging.root.addHandler(handler)
|
||||||
|
|
||||||
|
if msg_init is not None:
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
logger.info(msg_init)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# TODO make inf_utils.py
|
||||||
|
|
||||||
|
|
||||||
|
# region Gradual Latent hires fix
|
||||||
|
|
||||||
|
|
||||||
|
class GradualLatent:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
ratio,
|
||||||
|
start_timesteps,
|
||||||
|
every_n_steps,
|
||||||
|
ratio_step,
|
||||||
|
s_noise=1.0,
|
||||||
|
gaussian_blur_ksize=None,
|
||||||
|
gaussian_blur_sigma=0.5,
|
||||||
|
gaussian_blur_strength=0.5,
|
||||||
|
unsharp_target_x=True,
|
||||||
|
):
|
||||||
|
self.ratio = ratio
|
||||||
|
self.start_timesteps = start_timesteps
|
||||||
|
self.every_n_steps = every_n_steps
|
||||||
|
self.ratio_step = ratio_step
|
||||||
|
self.s_noise = s_noise
|
||||||
|
self.gaussian_blur_ksize = gaussian_blur_ksize
|
||||||
|
self.gaussian_blur_sigma = gaussian_blur_sigma
|
||||||
|
self.gaussian_blur_strength = gaussian_blur_strength
|
||||||
|
self.unsharp_target_x = unsharp_target_x
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return (
|
||||||
|
f"GradualLatent(ratio={self.ratio}, start_timesteps={self.start_timesteps}, "
|
||||||
|
+ f"every_n_steps={self.every_n_steps}, ratio_step={self.ratio_step}, s_noise={self.s_noise}, "
|
||||||
|
+ f"gaussian_blur_ksize={self.gaussian_blur_ksize}, gaussian_blur_sigma={self.gaussian_blur_sigma}, gaussian_blur_strength={self.gaussian_blur_strength}, "
|
||||||
|
+ f"unsharp_target_x={self.unsharp_target_x})"
|
||||||
|
)
|
||||||
|
|
||||||
|
def apply_unshark_mask(self, x: torch.Tensor):
|
||||||
|
if self.gaussian_blur_ksize is None:
|
||||||
|
return x
|
||||||
|
blurred = transforms.functional.gaussian_blur(x, self.gaussian_blur_ksize, self.gaussian_blur_sigma)
|
||||||
|
# mask = torch.sigmoid((x - blurred) * self.gaussian_blur_strength)
|
||||||
|
mask = (x - blurred) * self.gaussian_blur_strength
|
||||||
|
sharpened = x + mask
|
||||||
|
return sharpened
|
||||||
|
|
||||||
|
def interpolate(self, x: torch.Tensor, resized_size, unsharp=True):
|
||||||
|
org_dtype = x.dtype
|
||||||
|
if org_dtype == torch.bfloat16:
|
||||||
|
x = x.float()
|
||||||
|
|
||||||
|
x = torch.nn.functional.interpolate(x, size=resized_size, mode="bicubic", align_corners=False).to(dtype=org_dtype)
|
||||||
|
|
||||||
|
# apply unsharp mask / アンシャープマスクを適用する
|
||||||
|
if unsharp and self.gaussian_blur_ksize:
|
||||||
|
x = self.apply_unshark_mask(x)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class EulerAncestralDiscreteSchedulerGL(EulerAncestralDiscreteScheduler):
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.resized_size = None
|
||||||
|
self.gradual_latent = None
|
||||||
|
|
||||||
|
def set_gradual_latent_params(self, size, gradual_latent: GradualLatent):
|
||||||
|
self.resized_size = size
|
||||||
|
self.gradual_latent = gradual_latent
|
||||||
|
|
||||||
|
def step(
|
||||||
|
self,
|
||||||
|
model_output: torch.FloatTensor,
|
||||||
|
timestep: Union[float, torch.FloatTensor],
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
generator: Optional[torch.Generator] = None,
|
||||||
|
return_dict: bool = True,
|
||||||
|
) -> Union[EulerAncestralDiscreteSchedulerOutput, Tuple]:
|
||||||
|
"""
|
||||||
|
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||||
|
process from the learned model outputs (most often the predicted noise).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_output (`torch.FloatTensor`):
|
||||||
|
The direct output from learned diffusion model.
|
||||||
|
timestep (`float`):
|
||||||
|
The current discrete timestep in the diffusion chain.
|
||||||
|
sample (`torch.FloatTensor`):
|
||||||
|
A current instance of a sample created by the diffusion process.
|
||||||
|
generator (`torch.Generator`, *optional*):
|
||||||
|
A random number generator.
|
||||||
|
return_dict (`bool`):
|
||||||
|
Whether or not to return a
|
||||||
|
[`~schedulers.scheduling_euler_ancestral_discrete.EulerAncestralDiscreteSchedulerOutput`] or tuple.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~schedulers.scheduling_euler_ancestral_discrete.EulerAncestralDiscreteSchedulerOutput`] or `tuple`:
|
||||||
|
If return_dict is `True`,
|
||||||
|
[`~schedulers.scheduling_euler_ancestral_discrete.EulerAncestralDiscreteSchedulerOutput`] is returned,
|
||||||
|
otherwise a tuple is returned where the first element is the sample tensor.
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
if isinstance(timestep, int) or isinstance(timestep, torch.IntTensor) or isinstance(timestep, torch.LongTensor):
|
||||||
|
raise ValueError(
|
||||||
|
(
|
||||||
|
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||||
|
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||||
|
" one of the `scheduler.timesteps` as a timestep."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self.is_scale_input_called:
|
||||||
|
# logger.warning(
|
||||||
|
print(
|
||||||
|
"The `scale_model_input` function should be called before `step` to ensure correct denoising. "
|
||||||
|
"See `StableDiffusionPipeline` for a usage example."
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.step_index is None:
|
||||||
|
self._init_step_index(timestep)
|
||||||
|
|
||||||
|
sigma = self.sigmas[self.step_index]
|
||||||
|
|
||||||
|
# 1. compute predicted original sample (x_0) from sigma-scaled predicted noise
|
||||||
|
if self.config.prediction_type == "epsilon":
|
||||||
|
pred_original_sample = sample - sigma * model_output
|
||||||
|
elif self.config.prediction_type == "v_prediction":
|
||||||
|
# * c_out + input * c_skip
|
||||||
|
pred_original_sample = model_output * (-sigma / (sigma**2 + 1) ** 0.5) + (sample / (sigma**2 + 1))
|
||||||
|
elif self.config.prediction_type == "sample":
|
||||||
|
raise NotImplementedError("prediction_type not implemented yet: sample")
|
||||||
|
else:
|
||||||
|
raise ValueError(f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, or `v_prediction`")
|
||||||
|
|
||||||
|
sigma_from = self.sigmas[self.step_index]
|
||||||
|
sigma_to = self.sigmas[self.step_index + 1]
|
||||||
|
sigma_up = (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5
|
||||||
|
sigma_down = (sigma_to**2 - sigma_up**2) ** 0.5
|
||||||
|
|
||||||
|
# 2. Convert to an ODE derivative
|
||||||
|
derivative = (sample - pred_original_sample) / sigma
|
||||||
|
|
||||||
|
dt = sigma_down - sigma
|
||||||
|
|
||||||
|
device = model_output.device
|
||||||
|
if self.resized_size is None:
|
||||||
|
prev_sample = sample + derivative * dt
|
||||||
|
|
||||||
|
noise = diffusers.schedulers.scheduling_euler_ancestral_discrete.randn_tensor(
|
||||||
|
model_output.shape, dtype=model_output.dtype, device=device, generator=generator
|
||||||
|
)
|
||||||
|
s_noise = 1.0
|
||||||
|
else:
|
||||||
|
print("resized_size", self.resized_size, "model_output.shape", model_output.shape, "sample.shape", sample.shape)
|
||||||
|
s_noise = self.gradual_latent.s_noise
|
||||||
|
|
||||||
|
if self.gradual_latent.unsharp_target_x:
|
||||||
|
prev_sample = sample + derivative * dt
|
||||||
|
prev_sample = self.gradual_latent.interpolate(prev_sample, self.resized_size)
|
||||||
|
else:
|
||||||
|
sample = self.gradual_latent.interpolate(sample, self.resized_size)
|
||||||
|
derivative = self.gradual_latent.interpolate(derivative, self.resized_size, unsharp=False)
|
||||||
|
prev_sample = sample + derivative * dt
|
||||||
|
|
||||||
|
noise = diffusers.schedulers.scheduling_euler_ancestral_discrete.randn_tensor(
|
||||||
|
(model_output.shape[0], model_output.shape[1], self.resized_size[0], self.resized_size[1]),
|
||||||
|
dtype=model_output.dtype,
|
||||||
|
device=device,
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
|
||||||
|
prev_sample = prev_sample + noise * sigma_up * s_noise
|
||||||
|
|
||||||
|
# upon completion increase step index by one
|
||||||
|
self._step_index += 1
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (prev_sample,)
|
||||||
|
|
||||||
|
return EulerAncestralDiscreteSchedulerOutput(prev_sample=prev_sample, pred_original_sample=pred_original_sample)
|
||||||
|
|
||||||
|
|
||||||
|
# endregion
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
from .utils import setup_logging
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def main(file):
|
||||||
|
logger.info(f"loading: {file}")
|
||||||
|
if os.path.splitext(file)[1] == ".safetensors":
|
||||||
|
sd = load_file(file)
|
||||||
|
else:
|
||||||
|
sd = torch.load(file, map_location="cpu")
|
||||||
|
|
||||||
|
values = []
|
||||||
|
|
||||||
|
keys = list(sd.keys())
|
||||||
|
for key in keys:
|
||||||
|
if "lora_up" in key or "lora_down" in key:
|
||||||
|
values.append((key, sd[key]))
|
||||||
|
print(f"number of LoRA modules: {len(values)}")
|
||||||
|
|
||||||
|
if args.show_all_keys:
|
||||||
|
for key in [k for k in keys if k not in values]:
|
||||||
|
values.append((key, sd[key]))
|
||||||
|
print(f"number of all modules: {len(values)}")
|
||||||
|
|
||||||
|
for key, value in values:
|
||||||
|
value = value.to(torch.float32)
|
||||||
|
print(f"{key},{str(tuple(value.size())).replace(', ', '-')},{torch.mean(torch.abs(value))},{torch.min(torch.abs(value))}")
|
||||||
|
|
||||||
|
|
||||||
|
def setup_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("file", type=str, help="model file to check / 重みを確認するモデルファイル")
|
||||||
|
parser.add_argument("-s", "--show_all_keys", action="store_true", help="show all keys / 全てのキーを表示する")
|
||||||
|
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = setup_parser()
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
main(args.file)
|
||||||
+1403
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,762 @@
|
|||||||
|
# temporary minimum implementation of LoRA
|
||||||
|
# FLUX doesn't have Conv2d, so we ignore it
|
||||||
|
# TODO commonize with the original 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 math
|
||||||
|
import os
|
||||||
|
from typing import Dict, List, Optional, Tuple, Type, Union
|
||||||
|
from diffusers import AutoencoderKL
|
||||||
|
from transformers import CLIPTextModel
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
import re
|
||||||
|
#from ..library.utils import setup_logging
|
||||||
|
|
||||||
|
#setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class LoRAModule(torch.nn.Module):
|
||||||
|
"""
|
||||||
|
replaces forward method of the original Linear, instead of replacing the original Linear module.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
lora_name,
|
||||||
|
org_module: torch.nn.Module,
|
||||||
|
multiplier=1.0,
|
||||||
|
lora_dim=4,
|
||||||
|
alpha=1,
|
||||||
|
dropout=None,
|
||||||
|
rank_dropout=None,
|
||||||
|
module_dropout=None,
|
||||||
|
):
|
||||||
|
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||||
|
super().__init__()
|
||||||
|
self.lora_name = lora_name
|
||||||
|
|
||||||
|
if org_module.__class__.__name__ == "Conv2d":
|
||||||
|
in_dim = org_module.in_channels
|
||||||
|
out_dim = org_module.out_channels
|
||||||
|
else:
|
||||||
|
in_dim = org_module.in_features
|
||||||
|
out_dim = org_module.out_features
|
||||||
|
|
||||||
|
self.lora_dim = lora_dim
|
||||||
|
|
||||||
|
if org_module.__class__.__name__ == "Conv2d":
|
||||||
|
kernel_size = org_module.kernel_size
|
||||||
|
stride = org_module.stride
|
||||||
|
padding = org_module.padding
|
||||||
|
self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
|
||||||
|
self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False)
|
||||||
|
else:
|
||||||
|
self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
|
||||||
|
self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False)
|
||||||
|
|
||||||
|
if type(alpha) == torch.Tensor:
|
||||||
|
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||||
|
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
|
||||||
|
self.scale = alpha / self.lora_dim
|
||||||
|
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
|
||||||
|
|
||||||
|
# same as microsoft's
|
||||||
|
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
|
||||||
|
torch.nn.init.zeros_(self.lora_up.weight)
|
||||||
|
|
||||||
|
self.multiplier = multiplier
|
||||||
|
self.org_module = org_module # remove in applying
|
||||||
|
self.dropout = dropout
|
||||||
|
self.rank_dropout = rank_dropout
|
||||||
|
self.module_dropout = module_dropout
|
||||||
|
|
||||||
|
def apply_to(self):
|
||||||
|
self.org_forward = self.org_module.forward
|
||||||
|
self.org_module.forward = self.forward
|
||||||
|
del self.org_module
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
org_forwarded = self.org_forward(x)
|
||||||
|
|
||||||
|
# module dropout
|
||||||
|
if self.module_dropout is not None and self.training:
|
||||||
|
if torch.rand(1) < self.module_dropout:
|
||||||
|
return org_forwarded
|
||||||
|
|
||||||
|
lx = self.lora_down(x)
|
||||||
|
|
||||||
|
# normal dropout
|
||||||
|
if self.dropout is not None and self.training:
|
||||||
|
lx = torch.nn.functional.dropout(lx, p=self.dropout)
|
||||||
|
|
||||||
|
# rank dropout
|
||||||
|
if self.rank_dropout is not None and self.training:
|
||||||
|
mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout
|
||||||
|
if len(lx.size()) == 3:
|
||||||
|
mask = mask.unsqueeze(1) # for Text Encoder
|
||||||
|
elif len(lx.size()) == 4:
|
||||||
|
mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d
|
||||||
|
lx = lx * mask
|
||||||
|
|
||||||
|
# scaling for rank dropout: treat as if the rank is changed
|
||||||
|
# maskから計算することも考えられるが、augmentation的な効果を期待してrank_dropoutを用いる
|
||||||
|
scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
|
||||||
|
else:
|
||||||
|
scale = self.scale
|
||||||
|
|
||||||
|
lx = self.lora_up(lx)
|
||||||
|
|
||||||
|
return org_forwarded + lx * self.multiplier * scale
|
||||||
|
|
||||||
|
|
||||||
|
class LoRAInfModule(LoRAModule):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
lora_name,
|
||||||
|
org_module: torch.nn.Module,
|
||||||
|
multiplier=1.0,
|
||||||
|
lora_dim=4,
|
||||||
|
alpha=1,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
# no dropout for inference
|
||||||
|
super().__init__(lora_name, org_module, multiplier, lora_dim, alpha)
|
||||||
|
|
||||||
|
self.org_module_ref = [org_module] # 後から参照できるように
|
||||||
|
self.enabled = True
|
||||||
|
self.network: LoRANetwork = None
|
||||||
|
|
||||||
|
def set_network(self, network):
|
||||||
|
self.network = network
|
||||||
|
|
||||||
|
# freezeしてマージする
|
||||||
|
def merge_to(self, sd, dtype, device):
|
||||||
|
# extract weight from org_module
|
||||||
|
org_sd = self.org_module.state_dict()
|
||||||
|
weight = org_sd["weight"]
|
||||||
|
org_dtype = weight.dtype
|
||||||
|
org_device = weight.device
|
||||||
|
weight = weight.to(torch.float) # calc in float
|
||||||
|
|
||||||
|
if dtype is None:
|
||||||
|
dtype = org_dtype
|
||||||
|
if device is None:
|
||||||
|
device = org_device
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# merge weight
|
||||||
|
if len(weight.size()) == 2:
|
||||||
|
# linear
|
||||||
|
weight = weight + self.multiplier * (up_weight @ down_weight) * self.scale
|
||||||
|
elif down_weight.size()[2:4] == (1, 1):
|
||||||
|
# conv2d 1x1
|
||||||
|
weight = (
|
||||||
|
weight
|
||||||
|
+ self.multiplier
|
||||||
|
* (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
|
||||||
|
* self.scale
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# conv2d 3x3
|
||||||
|
conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3)
|
||||||
|
# logger.info(conved.size(), weight.size(), module.stride, module.padding)
|
||||||
|
weight = weight + self.multiplier * conved * 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):
|
||||||
|
if multiplier is None:
|
||||||
|
multiplier = self.multiplier
|
||||||
|
|
||||||
|
# get up/down weight from module
|
||||||
|
up_weight = self.lora_up.weight.to(torch.float)
|
||||||
|
down_weight = self.lora_down.weight.to(torch.float)
|
||||||
|
|
||||||
|
# pre-calculated weight
|
||||||
|
if len(down_weight.size()) == 2:
|
||||||
|
# linear
|
||||||
|
weight = self.multiplier * (up_weight @ down_weight) * self.scale
|
||||||
|
elif down_weight.size()[2:4] == (1, 1):
|
||||||
|
# conv2d 1x1
|
||||||
|
weight = (
|
||||||
|
self.multiplier
|
||||||
|
* (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
|
||||||
|
* self.scale
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# conv2d 3x3
|
||||||
|
conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3)
|
||||||
|
weight = self.multiplier * conved * self.scale
|
||||||
|
|
||||||
|
return weight
|
||||||
|
|
||||||
|
def set_region(self, region):
|
||||||
|
self.region = region
|
||||||
|
self.region_mask = None
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
if not self.enabled:
|
||||||
|
return self.org_forward(x)
|
||||||
|
return self.default_forward(x)
|
||||||
|
|
||||||
|
|
||||||
|
def create_network(
|
||||||
|
multiplier: float,
|
||||||
|
network_dim: Optional[int],
|
||||||
|
network_alpha: Optional[float],
|
||||||
|
ae: AutoencoderKL,
|
||||||
|
text_encoders: List[CLIPTextModel],
|
||||||
|
flux,
|
||||||
|
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)
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# single or double blocks
|
||||||
|
train_blocks = kwargs.get("train_blocks", None) # None (default), "all" (same as None), "single", "double"
|
||||||
|
if train_blocks is not None:
|
||||||
|
assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
||||||
|
|
||||||
|
# すごく引数が多いな ( ^ω^)・・・
|
||||||
|
network = LoRANetwork(
|
||||||
|
text_encoders,
|
||||||
|
flux,
|
||||||
|
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,
|
||||||
|
train_blocks=train_blocks,
|
||||||
|
varbose=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
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, flux, 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
|
||||||
|
modules_dim = {}
|
||||||
|
modules_alpha = {}
|
||||||
|
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)
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
return network, weights_sd
|
||||||
|
|
||||||
|
|
||||||
|
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", "CLIPMLP"]
|
||||||
|
LORA_PREFIX_FLUX = "lora_unet" # make ComfyUI compatible
|
||||||
|
LORA_PREFIX_TEXT_ENCODER_CLIP = "lora_te1"
|
||||||
|
LORA_PREFIX_TEXT_ENCODER_T5 = "lora_te2"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
text_encoders: Union[List[CLIPTextModel], CLIPTextModel],
|
||||||
|
unet,
|
||||||
|
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,
|
||||||
|
train_blocks: Optional[str] = None,
|
||||||
|
varbose: 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.train_blocks = train_blocks if train_blocks is not None else "all"
|
||||||
|
|
||||||
|
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")
|
||||||
|
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}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# create module instances
|
||||||
|
def create_modules(
|
||||||
|
is_flux: bool, text_encoder_idx: Optional[int], root_module: torch.nn.Module, target_replace_modules: List[str]
|
||||||
|
) -> List[LoRAModule]:
|
||||||
|
prefix = (
|
||||||
|
self.LORA_PREFIX_FLUX
|
||||||
|
if is_flux
|
||||||
|
else (self.LORA_PREFIX_TEXT_ENCODER_CLIP if text_encoder_idx == 0 else self.LORA_PREFIX_TEXT_ENCODER_T5)
|
||||||
|
)
|
||||||
|
|
||||||
|
loras = []
|
||||||
|
skipped = []
|
||||||
|
for name, module in root_module.named_modules():
|
||||||
|
if module.__class__.__name__ in target_replace_modules:
|
||||||
|
for child_name, child_module in module.named_modules():
|
||||||
|
is_linear = child_module.__class__.__name__ == "Linear"
|
||||||
|
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 + "." + child_name
|
||||||
|
lora_name = lora_name.replace(".", "_")
|
||||||
|
|
||||||
|
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 = self.lora_dim
|
||||||
|
alpha = self.alpha
|
||||||
|
elif self.conv_lora_dim is not None:
|
||||||
|
dim = self.conv_lora_dim
|
||||||
|
alpha = self.conv_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
|
||||||
|
|
||||||
|
lora = module_class(
|
||||||
|
lora_name,
|
||||||
|
child_module,
|
||||||
|
self.multiplier,
|
||||||
|
dim,
|
||||||
|
alpha,
|
||||||
|
dropout=dropout,
|
||||||
|
rank_dropout=rank_dropout,
|
||||||
|
module_dropout=module_dropout,
|
||||||
|
)
|
||||||
|
loras.append(lora)
|
||||||
|
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
|
||||||
|
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)
|
||||||
|
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":
|
||||||
|
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_DOUBLE + LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||||
|
elif self.train_blocks == "single":
|
||||||
|
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||||
|
elif self.train_blocks == "double":
|
||||||
|
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_DOUBLE
|
||||||
|
|
||||||
|
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.")
|
||||||
|
|
||||||
|
skipped = skipped_te + skipped_un
|
||||||
|
if varbose 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 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")
|
||||||
|
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, flux, 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) or key.startswith(LoRANetwork.LORA_PREFIX_TEXT_ENCODER_T5):
|
||||||
|
apply_text_encoder = True
|
||||||
|
elif key.startswith(LoRANetwork.LORA_PREFIX_FLUX):
|
||||||
|
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}")
|
||||||
|
|
||||||
|
# 二つの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の組み合わせはサポートされていません"
|
||||||
|
|
||||||
|
self.requires_grad_(True)
|
||||||
|
|
||||||
|
all_params = []
|
||||||
|
lr_descriptions = []
|
||||||
|
|
||||||
|
def assemble_params(loras, lr, 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:
|
||||||
|
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 * 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:
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
all_params.extend(params)
|
||||||
|
lr_descriptions.extend(["textencoder" + (" " + 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,
|
||||||
|
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)
|
||||||
@@ -0,0 +1,360 @@
|
|||||||
|
import math
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file, save_file
|
||||||
|
from library import sai_model_spec, train_util
|
||||||
|
import library.model_util as model_util
|
||||||
|
import lora
|
||||||
|
from .utils import setup_logging
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def load_state_dict(file_name, dtype):
|
||||||
|
if os.path.splitext(file_name)[1] == ".safetensors":
|
||||||
|
sd = load_file(file_name)
|
||||||
|
metadata = train_util.load_metadata_from_safetensors(file_name)
|
||||||
|
else:
|
||||||
|
sd = torch.load(file_name, map_location="cpu")
|
||||||
|
metadata = {}
|
||||||
|
|
||||||
|
for key in list(sd.keys()):
|
||||||
|
if type(sd[key]) == torch.Tensor:
|
||||||
|
sd[key] = sd[key].to(dtype)
|
||||||
|
|
||||||
|
return sd, metadata
|
||||||
|
|
||||||
|
|
||||||
|
def save_to_file(file_name, model, state_dict, dtype, metadata):
|
||||||
|
if dtype is not None:
|
||||||
|
for key in list(state_dict.keys()):
|
||||||
|
if type(state_dict[key]) == torch.Tensor:
|
||||||
|
state_dict[key] = state_dict[key].to(dtype)
|
||||||
|
|
||||||
|
if os.path.splitext(file_name)[1] == ".safetensors":
|
||||||
|
save_file(model, file_name, metadata=metadata)
|
||||||
|
else:
|
||||||
|
torch.save(model, file_name)
|
||||||
|
|
||||||
|
|
||||||
|
def merge_to_sd_model(text_encoder, unet, models, ratios, merge_dtype):
|
||||||
|
text_encoder.to(merge_dtype)
|
||||||
|
unet.to(merge_dtype)
|
||||||
|
|
||||||
|
# create module map
|
||||||
|
name_to_module = {}
|
||||||
|
for i, root_module in enumerate([text_encoder, unet]):
|
||||||
|
if i == 0:
|
||||||
|
prefix = lora.LoRANetwork.LORA_PREFIX_TEXT_ENCODER
|
||||||
|
target_replace_modules = lora.LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE
|
||||||
|
else:
|
||||||
|
prefix = lora.LoRANetwork.LORA_PREFIX_UNET
|
||||||
|
target_replace_modules = (
|
||||||
|
lora.LoRANetwork.UNET_TARGET_REPLACE_MODULE + lora.LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, module in root_module.named_modules():
|
||||||
|
if module.__class__.__name__ in target_replace_modules:
|
||||||
|
for child_name, child_module in module.named_modules():
|
||||||
|
if child_module.__class__.__name__ == "Linear" or child_module.__class__.__name__ == "Conv2d":
|
||||||
|
lora_name = prefix + "." + name + "." + child_name
|
||||||
|
lora_name = lora_name.replace(".", "_")
|
||||||
|
name_to_module[lora_name] = child_module
|
||||||
|
|
||||||
|
for model, ratio in zip(models, ratios):
|
||||||
|
logger.info(f"loading: {model}")
|
||||||
|
lora_sd, _ = load_state_dict(model, merge_dtype)
|
||||||
|
|
||||||
|
logger.info(f"merging...")
|
||||||
|
for key in lora_sd.keys():
|
||||||
|
if "lora_down" in key:
|
||||||
|
up_key = key.replace("lora_down", "lora_up")
|
||||||
|
alpha_key = key[: key.index("lora_down")] + "alpha"
|
||||||
|
|
||||||
|
# find original module for this lora
|
||||||
|
module_name = ".".join(key.split(".")[:-2]) # remove trailing ".lora_down.weight"
|
||||||
|
if module_name not in name_to_module:
|
||||||
|
logger.info(f"no module found for LoRA weight: {key}")
|
||||||
|
continue
|
||||||
|
module = name_to_module[module_name]
|
||||||
|
# logger.info(f"apply {key} to {module}")
|
||||||
|
|
||||||
|
down_weight = lora_sd[key]
|
||||||
|
up_weight = lora_sd[up_key]
|
||||||
|
|
||||||
|
dim = down_weight.size()[0]
|
||||||
|
alpha = lora_sd.get(alpha_key, dim)
|
||||||
|
scale = alpha / dim
|
||||||
|
|
||||||
|
# W <- W + U * D
|
||||||
|
weight = module.weight
|
||||||
|
if len(weight.size()) == 2:
|
||||||
|
# linear
|
||||||
|
if len(up_weight.size()) == 4: # use linear projection mismatch
|
||||||
|
up_weight = up_weight.squeeze(3).squeeze(2)
|
||||||
|
down_weight = down_weight.squeeze(3).squeeze(2)
|
||||||
|
weight = weight + ratio * (up_weight @ down_weight) * scale
|
||||||
|
elif down_weight.size()[2:4] == (1, 1):
|
||||||
|
# conv2d 1x1
|
||||||
|
weight = (
|
||||||
|
weight
|
||||||
|
+ ratio
|
||||||
|
* (up_weight.squeeze(3).squeeze(2) @ down_weight.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3)
|
||||||
|
* scale
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# conv2d 3x3
|
||||||
|
conved = torch.nn.functional.conv2d(down_weight.permute(1, 0, 2, 3), up_weight).permute(1, 0, 2, 3)
|
||||||
|
# logger.info(conved.size(), weight.size(), module.stride, module.padding)
|
||||||
|
weight = weight + ratio * conved * scale
|
||||||
|
|
||||||
|
module.weight = torch.nn.Parameter(weight)
|
||||||
|
|
||||||
|
|
||||||
|
def merge_lora_models(models, ratios, merge_dtype, concat=False, shuffle=False):
|
||||||
|
base_alphas = {} # alpha for merged model
|
||||||
|
base_dims = {}
|
||||||
|
|
||||||
|
merged_sd = {}
|
||||||
|
v2 = None
|
||||||
|
base_model = None
|
||||||
|
for model, ratio in zip(models, ratios):
|
||||||
|
logger.info(f"loading: {model}")
|
||||||
|
lora_sd, lora_metadata = load_state_dict(model, merge_dtype)
|
||||||
|
|
||||||
|
if lora_metadata is not None:
|
||||||
|
if v2 is None:
|
||||||
|
v2 = lora_metadata.get(train_util.SS_METADATA_KEY_V2, None) # return string
|
||||||
|
if base_model is None:
|
||||||
|
base_model = lora_metadata.get(train_util.SS_METADATA_KEY_BASE_MODEL_VERSION, None)
|
||||||
|
|
||||||
|
# get alpha and dim
|
||||||
|
alphas = {} # alpha for current model
|
||||||
|
dims = {} # dims for current model
|
||||||
|
for key in lora_sd.keys():
|
||||||
|
if "alpha" in key:
|
||||||
|
lora_module_name = key[: key.rfind(".alpha")]
|
||||||
|
alpha = float(lora_sd[key].detach().numpy())
|
||||||
|
alphas[lora_module_name] = alpha
|
||||||
|
if lora_module_name not in base_alphas:
|
||||||
|
base_alphas[lora_module_name] = alpha
|
||||||
|
elif "lora_down" in key:
|
||||||
|
lora_module_name = key[: key.rfind(".lora_down")]
|
||||||
|
dim = lora_sd[key].size()[0]
|
||||||
|
dims[lora_module_name] = dim
|
||||||
|
if lora_module_name not in base_dims:
|
||||||
|
base_dims[lora_module_name] = dim
|
||||||
|
|
||||||
|
for lora_module_name in dims.keys():
|
||||||
|
if lora_module_name not in alphas:
|
||||||
|
alpha = dims[lora_module_name]
|
||||||
|
alphas[lora_module_name] = alpha
|
||||||
|
if lora_module_name not in base_alphas:
|
||||||
|
base_alphas[lora_module_name] = alpha
|
||||||
|
|
||||||
|
logger.info(f"dim: {list(set(dims.values()))}, alpha: {list(set(alphas.values()))}")
|
||||||
|
|
||||||
|
# merge
|
||||||
|
logger.info(f"merging...")
|
||||||
|
for key in lora_sd.keys():
|
||||||
|
if "alpha" in key:
|
||||||
|
continue
|
||||||
|
if "lora_up" in key and concat:
|
||||||
|
concat_dim = 1
|
||||||
|
elif "lora_down" in key and concat:
|
||||||
|
concat_dim = 0
|
||||||
|
else:
|
||||||
|
concat_dim = None
|
||||||
|
|
||||||
|
lora_module_name = key[: key.rfind(".lora_")]
|
||||||
|
|
||||||
|
base_alpha = base_alphas[lora_module_name]
|
||||||
|
alpha = alphas[lora_module_name]
|
||||||
|
|
||||||
|
scale = math.sqrt(alpha / base_alpha) * ratio
|
||||||
|
scale = abs(scale) if "lora_up" in key else scale # マイナスの重みに対応する。
|
||||||
|
|
||||||
|
if key in merged_sd:
|
||||||
|
assert (
|
||||||
|
merged_sd[key].size() == lora_sd[key].size() or concat_dim is not None
|
||||||
|
), f"weights shape mismatch merging v1 and v2, different dims? / 重みのサイズが合いません。v1とv2、または次元数の異なるモデルはマージできません"
|
||||||
|
if concat_dim is not None:
|
||||||
|
merged_sd[key] = torch.cat([merged_sd[key], lora_sd[key] * scale], dim=concat_dim)
|
||||||
|
else:
|
||||||
|
merged_sd[key] = merged_sd[key] + lora_sd[key] * scale
|
||||||
|
else:
|
||||||
|
merged_sd[key] = lora_sd[key] * scale
|
||||||
|
|
||||||
|
# set alpha to sd
|
||||||
|
for lora_module_name, alpha in base_alphas.items():
|
||||||
|
key = lora_module_name + ".alpha"
|
||||||
|
merged_sd[key] = torch.tensor(alpha)
|
||||||
|
if shuffle:
|
||||||
|
key_down = lora_module_name + ".lora_down.weight"
|
||||||
|
key_up = lora_module_name + ".lora_up.weight"
|
||||||
|
dim = merged_sd[key_down].shape[0]
|
||||||
|
perm = torch.randperm(dim)
|
||||||
|
merged_sd[key_down] = merged_sd[key_down][perm]
|
||||||
|
merged_sd[key_up] = merged_sd[key_up][:,perm]
|
||||||
|
|
||||||
|
logger.info("merged model")
|
||||||
|
logger.info(f"dim: {list(set(base_dims.values()))}, alpha: {list(set(base_alphas.values()))}")
|
||||||
|
|
||||||
|
# check all dims are same
|
||||||
|
dims_list = list(set(base_dims.values()))
|
||||||
|
alphas_list = list(set(base_alphas.values()))
|
||||||
|
all_same_dims = True
|
||||||
|
all_same_alphas = True
|
||||||
|
for dims in dims_list:
|
||||||
|
if dims != dims_list[0]:
|
||||||
|
all_same_dims = False
|
||||||
|
break
|
||||||
|
for alphas in alphas_list:
|
||||||
|
if alphas != alphas_list[0]:
|
||||||
|
all_same_alphas = False
|
||||||
|
break
|
||||||
|
|
||||||
|
# build minimum metadata
|
||||||
|
dims = f"{dims_list[0]}" if all_same_dims else "Dynamic"
|
||||||
|
alphas = f"{alphas_list[0]}" if all_same_alphas else "Dynamic"
|
||||||
|
metadata = train_util.build_minimum_network_metadata(v2, base_model, "networks.lora", dims, alphas, None)
|
||||||
|
|
||||||
|
return merged_sd, metadata, v2 == "True"
|
||||||
|
|
||||||
|
|
||||||
|
def merge(args):
|
||||||
|
assert len(args.models) == len(args.ratios), f"number of models must be equal to number of ratios / モデルの数と重みの数は合わせてください"
|
||||||
|
|
||||||
|
def str_to_dtype(p):
|
||||||
|
if p == "float":
|
||||||
|
return torch.float
|
||||||
|
if p == "fp16":
|
||||||
|
return torch.float16
|
||||||
|
if p == "bf16":
|
||||||
|
return torch.bfloat16
|
||||||
|
return None
|
||||||
|
|
||||||
|
merge_dtype = str_to_dtype(args.precision)
|
||||||
|
save_dtype = str_to_dtype(args.save_precision)
|
||||||
|
if save_dtype is None:
|
||||||
|
save_dtype = merge_dtype
|
||||||
|
|
||||||
|
if args.sd_model is not None:
|
||||||
|
logger.info(f"loading SD model: {args.sd_model}")
|
||||||
|
|
||||||
|
text_encoder, vae, unet = model_util.load_models_from_stable_diffusion_checkpoint(args.v2, args.sd_model)
|
||||||
|
|
||||||
|
merge_to_sd_model(text_encoder, unet, args.models, args.ratios, merge_dtype)
|
||||||
|
|
||||||
|
if args.no_metadata:
|
||||||
|
sai_metadata = None
|
||||||
|
else:
|
||||||
|
merged_from = sai_model_spec.build_merged_from([args.sd_model] + args.models)
|
||||||
|
title = os.path.splitext(os.path.basename(args.save_to))[0]
|
||||||
|
sai_metadata = sai_model_spec.build_metadata(
|
||||||
|
None,
|
||||||
|
args.v2,
|
||||||
|
args.v2,
|
||||||
|
False,
|
||||||
|
False,
|
||||||
|
False,
|
||||||
|
time.time(),
|
||||||
|
title=title,
|
||||||
|
merged_from=merged_from,
|
||||||
|
is_stable_diffusion_ckpt=True,
|
||||||
|
)
|
||||||
|
if args.v2:
|
||||||
|
# TODO read sai modelspec
|
||||||
|
logger.warning(
|
||||||
|
"Cannot determine if model is for v-prediction, so save metadata as v-prediction / modelがv-prediction用か否か不明なため、仮にv-prediction用としてmetadataを保存します"
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(f"saving SD model to: {args.save_to}")
|
||||||
|
model_util.save_stable_diffusion_checkpoint(
|
||||||
|
args.v2, args.save_to, text_encoder, unet, args.sd_model, 0, 0, sai_metadata, save_dtype, vae
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
state_dict, metadata, v2 = merge_lora_models(args.models, args.ratios, merge_dtype, args.concat, args.shuffle)
|
||||||
|
|
||||||
|
logger.info(f"calculating hashes and creating 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
|
||||||
|
|
||||||
|
if not args.no_metadata:
|
||||||
|
merged_from = sai_model_spec.build_merged_from(args.models)
|
||||||
|
title = os.path.splitext(os.path.basename(args.save_to))[0]
|
||||||
|
sai_metadata = sai_model_spec.build_metadata(
|
||||||
|
state_dict, v2, v2, False, True, False, time.time(), title=title, merged_from=merged_from
|
||||||
|
)
|
||||||
|
if v2:
|
||||||
|
# TODO read sai modelspec
|
||||||
|
logger.warning(
|
||||||
|
"Cannot determine if LoRA is for v-prediction, so save metadata as v-prediction / LoRAがv-prediction用か否か不明なため、仮にv-prediction用としてmetadataを保存します"
|
||||||
|
)
|
||||||
|
metadata.update(sai_metadata)
|
||||||
|
|
||||||
|
logger.info(f"saving model to: {args.save_to}")
|
||||||
|
save_to_file(args.save_to, state_dict, state_dict, save_dtype, metadata)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--v2", action="store_true", help="load Stable Diffusion v2.x model / Stable Diffusion 2.xのモデルを読み込む")
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_precision",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=[None, "float", "fp16", "bf16"],
|
||||||
|
help="precision in saving, same to merging if omitted / 保存時に精度を変更して保存する、省略時はマージ時の精度と同じ",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--precision",
|
||||||
|
type=str,
|
||||||
|
default="float",
|
||||||
|
choices=["float", "fp16", "bf16"],
|
||||||
|
help="precision in merging (float is recommended) / マージの計算時の精度(floatを推奨)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--sd_model",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Stable Diffusion model to load: ckpt or safetensors file, merge LoRA models if omitted / 読み込むモデル、ckptまたはsafetensors。省略時はLoRAモデル同士をマージする",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_to", type=str, default=None, help="destination file name: ckpt or safetensors file / 保存先のファイル名、ckptまたはsafetensors"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--models", type=str, nargs="*", help="LoRA models to merge: ckpt or safetensors file / マージするLoRAモデル、ckptまたはsafetensors"
|
||||||
|
)
|
||||||
|
parser.add_argument("--ratios", type=float, nargs="*", help="ratios for each model / それぞれのLoRAモデルの比率")
|
||||||
|
parser.add_argument(
|
||||||
|
"--no_metadata",
|
||||||
|
action="store_true",
|
||||||
|
help="do not save sai modelspec metadata (minimum ss_metadata for LoRA is saved) / "
|
||||||
|
+ "sai modelspecのメタデータを保存しない(LoRAの最低限のss_metadataは保存される)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--concat",
|
||||||
|
action="store_true",
|
||||||
|
help="concat lora instead of merge (The dim(rank) of the output LoRA is the sum of the input dims) / "
|
||||||
|
+ "マージの代わりに結合する(LoRAのdim(rank)は入力dimの合計になる)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--shuffle",
|
||||||
|
action="store_true",
|
||||||
|
help="shuffle lora weight./ "
|
||||||
|
+ "LoRAの重みをシャッフルする",
|
||||||
|
)
|
||||||
|
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = setup_parser()
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
merge(args)
|
||||||
@@ -0,0 +1,411 @@
|
|||||||
|
# Convert LoRA to different rank approximation (should only be used to go to lower rank)
|
||||||
|
# This code is based off the extract_lora_from_models.py file which is based on https://github.com/cloneofsimo/lora/blob/develop/lora_diffusion/cli_svd.py
|
||||||
|
# Thanks to cloneofsimo
|
||||||
|
|
||||||
|
import os
|
||||||
|
import argparse
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file, save_file, safe_open
|
||||||
|
from tqdm import tqdm
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from library import train_util
|
||||||
|
from library import model_util
|
||||||
|
from .utils import setup_logging
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
MIN_SV = 1e-6
|
||||||
|
|
||||||
|
# Model save and load functions
|
||||||
|
|
||||||
|
|
||||||
|
def load_state_dict(file_name, dtype):
|
||||||
|
if model_util.is_safetensors(file_name):
|
||||||
|
sd = load_file(file_name)
|
||||||
|
with safe_open(file_name, framework="pt") as f:
|
||||||
|
metadata = f.metadata()
|
||||||
|
else:
|
||||||
|
sd = torch.load(file_name, map_location="cpu")
|
||||||
|
metadata = None
|
||||||
|
|
||||||
|
for key in list(sd.keys()):
|
||||||
|
if type(sd[key]) == torch.Tensor:
|
||||||
|
sd[key] = sd[key].to(dtype)
|
||||||
|
|
||||||
|
return sd, metadata
|
||||||
|
|
||||||
|
|
||||||
|
def save_to_file(file_name, state_dict, dtype, metadata):
|
||||||
|
if dtype is not None:
|
||||||
|
for key in list(state_dict.keys()):
|
||||||
|
if type(state_dict[key]) == torch.Tensor:
|
||||||
|
state_dict[key] = state_dict[key].to(dtype)
|
||||||
|
|
||||||
|
if model_util.is_safetensors(file_name):
|
||||||
|
save_file(state_dict, file_name, metadata)
|
||||||
|
else:
|
||||||
|
torch.save(state_dict, file_name)
|
||||||
|
|
||||||
|
|
||||||
|
# Indexing functions
|
||||||
|
|
||||||
|
|
||||||
|
def index_sv_cumulative(S, target):
|
||||||
|
original_sum = float(torch.sum(S))
|
||||||
|
cumulative_sums = torch.cumsum(S, dim=0) / original_sum
|
||||||
|
index = int(torch.searchsorted(cumulative_sums, target)) + 1
|
||||||
|
index = max(1, min(index, len(S) - 1))
|
||||||
|
|
||||||
|
return index
|
||||||
|
|
||||||
|
|
||||||
|
def index_sv_fro(S, target):
|
||||||
|
S_squared = S.pow(2)
|
||||||
|
S_fro_sq = float(torch.sum(S_squared))
|
||||||
|
sum_S_squared = torch.cumsum(S_squared, dim=0) / S_fro_sq
|
||||||
|
index = int(torch.searchsorted(sum_S_squared, target**2)) + 1
|
||||||
|
index = max(1, min(index, len(S) - 1))
|
||||||
|
|
||||||
|
return index
|
||||||
|
|
||||||
|
|
||||||
|
def index_sv_ratio(S, target):
|
||||||
|
max_sv = S[0]
|
||||||
|
min_sv = max_sv / target
|
||||||
|
index = int(torch.sum(S > min_sv).item())
|
||||||
|
index = max(1, min(index, len(S) - 1))
|
||||||
|
|
||||||
|
return index
|
||||||
|
|
||||||
|
|
||||||
|
# Modified from Kohaku-blueleaf's extract/merge functions
|
||||||
|
def extract_conv(weight, lora_rank, dynamic_method, dynamic_param, device, scale=1):
|
||||||
|
out_size, in_size, kernel_size, _ = weight.size()
|
||||||
|
U, S, Vh = torch.linalg.svd(weight.reshape(out_size, -1).to(device))
|
||||||
|
|
||||||
|
param_dict = rank_resize(S, lora_rank, dynamic_method, dynamic_param, scale)
|
||||||
|
lora_rank = param_dict["new_rank"]
|
||||||
|
|
||||||
|
U = U[:, :lora_rank]
|
||||||
|
S = S[:lora_rank]
|
||||||
|
U = U @ torch.diag(S)
|
||||||
|
Vh = Vh[:lora_rank, :]
|
||||||
|
|
||||||
|
param_dict["lora_down"] = Vh.reshape(lora_rank, in_size, kernel_size, kernel_size).cpu()
|
||||||
|
param_dict["lora_up"] = U.reshape(out_size, lora_rank, 1, 1).cpu()
|
||||||
|
del U, S, Vh, weight
|
||||||
|
return param_dict
|
||||||
|
|
||||||
|
|
||||||
|
def extract_linear(weight, lora_rank, dynamic_method, dynamic_param, device, scale=1):
|
||||||
|
out_size, in_size = weight.size()
|
||||||
|
|
||||||
|
U, S, Vh = torch.linalg.svd(weight.to(device))
|
||||||
|
|
||||||
|
param_dict = rank_resize(S, lora_rank, dynamic_method, dynamic_param, scale)
|
||||||
|
lora_rank = param_dict["new_rank"]
|
||||||
|
|
||||||
|
U = U[:, :lora_rank]
|
||||||
|
S = S[:lora_rank]
|
||||||
|
U = U @ torch.diag(S)
|
||||||
|
Vh = Vh[:lora_rank, :]
|
||||||
|
|
||||||
|
param_dict["lora_down"] = Vh.reshape(lora_rank, in_size).cpu()
|
||||||
|
param_dict["lora_up"] = U.reshape(out_size, lora_rank).cpu()
|
||||||
|
del U, S, Vh, weight
|
||||||
|
return param_dict
|
||||||
|
|
||||||
|
|
||||||
|
def merge_conv(lora_down, lora_up, device):
|
||||||
|
in_rank, in_size, kernel_size, k_ = lora_down.shape
|
||||||
|
out_size, out_rank, _, _ = lora_up.shape
|
||||||
|
assert in_rank == out_rank and kernel_size == k_, f"rank {in_rank} {out_rank} or kernel {kernel_size} {k_} mismatch"
|
||||||
|
|
||||||
|
lora_down = lora_down.to(device)
|
||||||
|
lora_up = lora_up.to(device)
|
||||||
|
|
||||||
|
merged = lora_up.reshape(out_size, -1) @ lora_down.reshape(in_rank, -1)
|
||||||
|
weight = merged.reshape(out_size, in_size, kernel_size, kernel_size)
|
||||||
|
del lora_up, lora_down
|
||||||
|
return weight
|
||||||
|
|
||||||
|
|
||||||
|
def merge_linear(lora_down, lora_up, device):
|
||||||
|
in_rank, in_size = lora_down.shape
|
||||||
|
out_size, out_rank = lora_up.shape
|
||||||
|
assert in_rank == out_rank, f"rank {in_rank} {out_rank} mismatch"
|
||||||
|
|
||||||
|
lora_down = lora_down.to(device)
|
||||||
|
lora_up = lora_up.to(device)
|
||||||
|
|
||||||
|
weight = lora_up @ lora_down
|
||||||
|
del lora_up, lora_down
|
||||||
|
return weight
|
||||||
|
|
||||||
|
|
||||||
|
# Calculate new rank
|
||||||
|
|
||||||
|
|
||||||
|
def rank_resize(S, rank, dynamic_method, dynamic_param, scale=1):
|
||||||
|
param_dict = {}
|
||||||
|
|
||||||
|
if dynamic_method == "sv_ratio":
|
||||||
|
# Calculate new dim and alpha based off ratio
|
||||||
|
new_rank = index_sv_ratio(S, dynamic_param) + 1
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
|
||||||
|
elif dynamic_method == "sv_cumulative":
|
||||||
|
# Calculate new dim and alpha based off cumulative sum
|
||||||
|
new_rank = index_sv_cumulative(S, dynamic_param) + 1
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
|
||||||
|
elif dynamic_method == "sv_fro":
|
||||||
|
# Calculate new dim and alpha based off sqrt sum of squares
|
||||||
|
new_rank = index_sv_fro(S, dynamic_param) + 1
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
else:
|
||||||
|
new_rank = rank
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
|
||||||
|
if S[0] <= MIN_SV: # Zero matrix, set dim to 1
|
||||||
|
new_rank = 1
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
elif new_rank > rank: # cap max rank at rank
|
||||||
|
new_rank = rank
|
||||||
|
new_alpha = float(scale * new_rank)
|
||||||
|
|
||||||
|
# Calculate resize info
|
||||||
|
s_sum = torch.sum(torch.abs(S))
|
||||||
|
s_rank = torch.sum(torch.abs(S[:new_rank]))
|
||||||
|
|
||||||
|
S_squared = S.pow(2)
|
||||||
|
s_fro = torch.sqrt(torch.sum(S_squared))
|
||||||
|
s_red_fro = torch.sqrt(torch.sum(S_squared[:new_rank]))
|
||||||
|
fro_percent = float(s_red_fro / s_fro)
|
||||||
|
|
||||||
|
param_dict["new_rank"] = new_rank
|
||||||
|
param_dict["new_alpha"] = new_alpha
|
||||||
|
param_dict["sum_retained"] = (s_rank) / s_sum
|
||||||
|
param_dict["fro_retained"] = fro_percent
|
||||||
|
param_dict["max_ratio"] = S[0] / S[new_rank - 1]
|
||||||
|
|
||||||
|
return param_dict
|
||||||
|
|
||||||
|
|
||||||
|
def resize_lora_model(lora_sd, new_rank, new_conv_rank, save_dtype, device, dynamic_method, dynamic_param, verbose):
|
||||||
|
network_alpha = None
|
||||||
|
network_dim = None
|
||||||
|
verbose_str = "\n"
|
||||||
|
fro_list = []
|
||||||
|
|
||||||
|
# Extract loaded lora dim and alpha
|
||||||
|
for key, value in lora_sd.items():
|
||||||
|
if network_alpha is None and "alpha" in key:
|
||||||
|
network_alpha = value
|
||||||
|
if network_dim is None and "lora_down" in key and len(value.size()) == 2:
|
||||||
|
network_dim = value.size()[0]
|
||||||
|
if network_alpha is not None and network_dim is not None:
|
||||||
|
break
|
||||||
|
if network_alpha is None:
|
||||||
|
network_alpha = network_dim
|
||||||
|
|
||||||
|
scale = network_alpha / network_dim
|
||||||
|
|
||||||
|
if dynamic_method:
|
||||||
|
logger.info(
|
||||||
|
f"Dynamically determining new alphas and dims based off {dynamic_method}: {dynamic_param}, max rank is {new_rank}"
|
||||||
|
)
|
||||||
|
|
||||||
|
lora_down_weight = None
|
||||||
|
lora_up_weight = None
|
||||||
|
|
||||||
|
o_lora_sd = lora_sd.copy()
|
||||||
|
block_down_name = None
|
||||||
|
block_up_name = None
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for key, value in tqdm(lora_sd.items()):
|
||||||
|
weight_name = None
|
||||||
|
if "lora_down" in key:
|
||||||
|
block_down_name = key.rsplit(".lora_down", 1)[0]
|
||||||
|
weight_name = key.rsplit(".", 1)[-1]
|
||||||
|
lora_down_weight = value
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# find corresponding lora_up and alpha
|
||||||
|
block_up_name = block_down_name
|
||||||
|
lora_up_weight = lora_sd.get(block_up_name + ".lora_up." + weight_name, None)
|
||||||
|
lora_alpha = lora_sd.get(block_down_name + ".alpha", None)
|
||||||
|
|
||||||
|
weights_loaded = lora_down_weight is not None and lora_up_weight is not None
|
||||||
|
|
||||||
|
if weights_loaded:
|
||||||
|
|
||||||
|
conv2d = len(lora_down_weight.size()) == 4
|
||||||
|
if lora_alpha is None:
|
||||||
|
scale = 1.0
|
||||||
|
else:
|
||||||
|
scale = lora_alpha / lora_down_weight.size()[0]
|
||||||
|
|
||||||
|
if conv2d:
|
||||||
|
full_weight_matrix = merge_conv(lora_down_weight, lora_up_weight, device)
|
||||||
|
param_dict = extract_conv(full_weight_matrix, new_conv_rank, dynamic_method, dynamic_param, device, scale)
|
||||||
|
else:
|
||||||
|
full_weight_matrix = merge_linear(lora_down_weight, lora_up_weight, device)
|
||||||
|
param_dict = extract_linear(full_weight_matrix, new_rank, dynamic_method, dynamic_param, device, scale)
|
||||||
|
|
||||||
|
if verbose:
|
||||||
|
max_ratio = param_dict["max_ratio"]
|
||||||
|
sum_retained = param_dict["sum_retained"]
|
||||||
|
fro_retained = param_dict["fro_retained"]
|
||||||
|
if not np.isnan(fro_retained):
|
||||||
|
fro_list.append(float(fro_retained))
|
||||||
|
|
||||||
|
verbose_str += f"{block_down_name:75} | "
|
||||||
|
verbose_str += (
|
||||||
|
f"sum(S) retained: {sum_retained:.1%}, fro retained: {fro_retained:.1%}, max(S) ratio: {max_ratio:0.1f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if verbose and dynamic_method:
|
||||||
|
verbose_str += f", dynamic | dim: {param_dict['new_rank']}, alpha: {param_dict['new_alpha']}\n"
|
||||||
|
else:
|
||||||
|
verbose_str += "\n"
|
||||||
|
|
||||||
|
new_alpha = param_dict["new_alpha"]
|
||||||
|
o_lora_sd[block_down_name + "." + "lora_down.weight"] = param_dict["lora_down"].to(save_dtype).contiguous()
|
||||||
|
o_lora_sd[block_up_name + "." + "lora_up.weight"] = param_dict["lora_up"].to(save_dtype).contiguous()
|
||||||
|
o_lora_sd[block_up_name + "." "alpha"] = torch.tensor(param_dict["new_alpha"]).to(save_dtype)
|
||||||
|
|
||||||
|
block_down_name = None
|
||||||
|
block_up_name = None
|
||||||
|
lora_down_weight = None
|
||||||
|
lora_up_weight = None
|
||||||
|
weights_loaded = False
|
||||||
|
del param_dict
|
||||||
|
|
||||||
|
if verbose:
|
||||||
|
print(verbose_str)
|
||||||
|
print(f"Average Frobenius norm retention: {np.mean(fro_list):.2%} | std: {np.std(fro_list):0.3f}")
|
||||||
|
logger.info("resizing complete")
|
||||||
|
return o_lora_sd, network_dim, new_alpha
|
||||||
|
|
||||||
|
|
||||||
|
def resize(args):
|
||||||
|
if args.save_to is None or not (
|
||||||
|
args.save_to.endswith(".ckpt")
|
||||||
|
or args.save_to.endswith(".pt")
|
||||||
|
or args.save_to.endswith(".pth")
|
||||||
|
or args.save_to.endswith(".safetensors")
|
||||||
|
):
|
||||||
|
raise Exception("The --save_to argument must be specified and must be a .ckpt , .pt, .pth or .safetensors file.")
|
||||||
|
|
||||||
|
args.new_conv_rank = args.new_conv_rank if args.new_conv_rank is not None else args.new_rank
|
||||||
|
|
||||||
|
def str_to_dtype(p):
|
||||||
|
if p == "float":
|
||||||
|
return torch.float
|
||||||
|
if p == "fp16":
|
||||||
|
return torch.float16
|
||||||
|
if p == "bf16":
|
||||||
|
return torch.bfloat16
|
||||||
|
return None
|
||||||
|
|
||||||
|
if args.dynamic_method and not args.dynamic_param:
|
||||||
|
raise Exception("If using dynamic_method, then dynamic_param is required")
|
||||||
|
|
||||||
|
merge_dtype = str_to_dtype("float") # matmul method above only seems to work in float32
|
||||||
|
save_dtype = str_to_dtype(args.save_precision)
|
||||||
|
if save_dtype is None:
|
||||||
|
save_dtype = merge_dtype
|
||||||
|
|
||||||
|
logger.info("loading Model...")
|
||||||
|
lora_sd, metadata = load_state_dict(args.model, merge_dtype)
|
||||||
|
|
||||||
|
logger.info("Resizing Lora...")
|
||||||
|
state_dict, old_dim, new_alpha = resize_lora_model(
|
||||||
|
lora_sd, args.new_rank, args.new_conv_rank, save_dtype, args.device, args.dynamic_method, args.dynamic_param, args.verbose
|
||||||
|
)
|
||||||
|
|
||||||
|
# update metadata
|
||||||
|
if metadata is None:
|
||||||
|
metadata = {}
|
||||||
|
|
||||||
|
comment = metadata.get("ss_training_comment", "")
|
||||||
|
|
||||||
|
if not args.dynamic_method:
|
||||||
|
conv_desc = "" if args.new_rank == args.new_conv_rank else f" (conv: {args.new_conv_rank})"
|
||||||
|
metadata["ss_training_comment"] = f"dimension is resized from {old_dim} to {args.new_rank}{conv_desc}; {comment}"
|
||||||
|
metadata["ss_network_dim"] = str(args.new_rank)
|
||||||
|
metadata["ss_network_alpha"] = str(new_alpha)
|
||||||
|
else:
|
||||||
|
metadata["ss_training_comment"] = (
|
||||||
|
f"Dynamic resize with {args.dynamic_method}: {args.dynamic_param} from {old_dim}; {comment}"
|
||||||
|
)
|
||||||
|
metadata["ss_network_dim"] = "Dynamic"
|
||||||
|
metadata["ss_network_alpha"] = "Dynamic"
|
||||||
|
|
||||||
|
model_hash, legacy_hash = train_util.precalculate_safetensors_hashes(state_dict, metadata)
|
||||||
|
metadata["sshs_model_hash"] = model_hash
|
||||||
|
metadata["sshs_legacy_hash"] = legacy_hash
|
||||||
|
|
||||||
|
logger.info(f"saving model to: {args.save_to}")
|
||||||
|
save_to_file(args.save_to, state_dict, save_dtype, metadata)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_precision",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=[None, "float", "fp16", "bf16"],
|
||||||
|
help="precision in saving, float if omitted / 保存時の精度、未指定時はfloat",
|
||||||
|
)
|
||||||
|
parser.add_argument("--new_rank", type=int, default=4, help="Specify rank of output LoRA / 出力するLoRAのrank (dim)")
|
||||||
|
parser.add_argument(
|
||||||
|
"--new_conv_rank",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Specify rank of output LoRA for Conv2d 3x3, None for same as new_rank / 出力するConv2D 3x3 LoRAのrank (dim)、Noneでnew_rankと同じ",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_to",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="destination file name: ckpt or safetensors file / 保存先のファイル名、ckptまたはsafetensors",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--model",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="LoRA model to resize at to new rank: ckpt or safetensors file / 読み込むLoRAモデル、ckptまたはsafetensors",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--device", type=str, default=None, help="device to use, cuda for GPU / 計算を行うデバイス、cuda でGPUを使う"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--verbose", action="store_true", help="Display verbose resizing information / rank変更時の詳細情報を出力する"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dynamic_method",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=[None, "sv_ratio", "sv_fro", "sv_cumulative"],
|
||||||
|
help="Specify dynamic resizing method, --new_rank is used as a hard limit for max rank",
|
||||||
|
)
|
||||||
|
parser.add_argument("--dynamic_param", type=float, default=None, help="Specify target for dynamic reduction")
|
||||||
|
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = setup_parser()
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
resize(args)
|
||||||
@@ -0,0 +1,764 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import math
|
||||||
|
import copy
|
||||||
|
import folder_paths
|
||||||
|
import comfy.model_management as mm
|
||||||
|
import comfy.utils
|
||||||
|
import argparse
|
||||||
|
from typing import Any, List
|
||||||
|
import time
|
||||||
|
|
||||||
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from accelerate import Accelerator
|
||||||
|
accelerator = Accelerator(mixed_precision='bf16', cpu=False)
|
||||||
|
from .library.device_utils import init_ipex, clean_memory_on_device
|
||||||
|
from .library.train_util import sample_images_common
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
from .library import flux_models, flux_train_utils, flux_utils, sd3_train_utils, strategy_base, strategy_flux, train_util
|
||||||
|
from .train_network import NetworkTrainer, setup_parser
|
||||||
|
|
||||||
|
import logging
|
||||||
|
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
class FluxNetworkTrainer(NetworkTrainer):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def assert_extra_args(self, args, train_dataset_group):
|
||||||
|
super().assert_extra_args(args, train_dataset_group)
|
||||||
|
|
||||||
|
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) # TODO check this
|
||||||
|
|
||||||
|
def load_target_model(self, args, weight_dtype, accelerator):
|
||||||
|
# currently offload to cpu for some models
|
||||||
|
name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" # TODO change this to a more robust way
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
clip_l = flux_utils.load_clip_l(args.clip_l, weight_dtype, "cpu")
|
||||||
|
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()
|
||||||
|
|
||||||
|
ae = flux_utils.load_ae(name, args.ae, weight_dtype, "cpu")
|
||||||
|
|
||||||
|
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):
|
||||||
|
return strategy_flux.FluxTokenizeStrategy(args.max_token_length, args.tokenizer_cache_dir)
|
||||||
|
|
||||||
|
def get_tokenizers(self, tokenize_strategy: strategy_flux.FluxTokenizeStrategy):
|
||||||
|
return [tokenize_strategy.clip_l, tokenize_strategy.t5xxl]
|
||||||
|
|
||||||
|
def get_latents_caching_strategy(self, args):
|
||||||
|
latents_caching_strategy = strategy_flux.FluxLatentsCachingStrategy(args.cache_latents_to_disk, args.vae_batch_size, False)
|
||||||
|
return latents_caching_strategy
|
||||||
|
|
||||||
|
def get_text_encoding_strategy(self, args):
|
||||||
|
return strategy_flux.FluxTextEncodingStrategy(apply_t5_attn_mask=args.apply_t5_attn_mask)
|
||||||
|
|
||||||
|
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_flux.FluxTextEncoderOutputsCachingStrategy(
|
||||||
|
args.cache_text_encoder_outputs_to_disk, None, False, 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)
|
||||||
|
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
|
||||||
|
with accelerator.autocast():
|
||||||
|
dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process)
|
||||||
|
|
||||||
|
# cache sample prompts
|
||||||
|
self.sample_prompts_te_outputs = None
|
||||||
|
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 = sd3_train_utils.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 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_t5_attn_mask
|
||||||
|
)
|
||||||
|
self.sample_prompts_te_outputs = sample_prompts_te_outputs
|
||||||
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
|
logger.info("move text encoders back to cpu")
|
||||||
|
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
|
||||||
|
text_encoders[0].to(accelerator.device, dtype=weight_dtype)
|
||||||
|
text_encoders[1].to(accelerator.device, dtype=weight_dtype)
|
||||||
|
|
||||||
|
def sample_images(self, accelerator, args, epoch, global_step, device, ae, tokenizer, text_encoder, flux):
|
||||||
|
if not args.split_mode:
|
||||||
|
flux_train_utils.sample_images(
|
||||||
|
accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
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):
|
||||||
|
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)
|
||||||
|
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)
|
||||||
|
|
||||||
|
wrapper = FluxUpperLowerWrapper(self.flux_upper, flux, accelerator.device)
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
flux_train_utils.sample_images(
|
||||||
|
accelerator, args, epoch, global_step, wrapper, ae, text_encoder, self.sample_prompts_te_outputs
|
||||||
|
)
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
|
||||||
|
def get_noise_scheduler(self, args: argparse.Namespace, device: torch.device) -> Any:
|
||||||
|
noise_scheduler = sd3_train_utils.FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=args.discrete_flow_shift)
|
||||||
|
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
|
||||||
|
|
||||||
|
def encode_images_to_latents(self, args, accelerator, vae, images):
|
||||||
|
return vae.encode(images).latent_dist.sample()
|
||||||
|
|
||||||
|
def shift_scale_latents(self, args, latents):
|
||||||
|
return 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,
|
||||||
|
):
|
||||||
|
# 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]
|
||||||
|
|
||||||
|
if args.timestep_sampling == "uniform" or args.timestep_sampling == "sigmoid":
|
||||||
|
# Simple random t-based noise sampling
|
||||||
|
if args.timestep_sampling == "sigmoid":
|
||||||
|
# https://github.com/XLabs-AI/x-flux/tree/main
|
||||||
|
t = torch.sigmoid(args.sigmoid_scale * torch.randn((bsz,), device=accelerator.device))
|
||||||
|
else:
|
||||||
|
t = torch.rand((bsz,), device=accelerator.device)
|
||||||
|
timesteps = t * 1000.0
|
||||||
|
t = t.view(-1, 1, 1, 1)
|
||||||
|
noisy_model_input = (1 - t) * latents + t * noise
|
||||||
|
else:
|
||||||
|
# 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,
|
||||||
|
)
|
||||||
|
indices = (u * self.noise_scheduler_copy.config.num_train_timesteps).long()
|
||||||
|
timesteps = self.noise_scheduler_copy.timesteps[indices].to(device=accelerator.device)
|
||||||
|
|
||||||
|
# Add noise according to flow matching.
|
||||||
|
sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=weight_dtype)
|
||||||
|
noisy_model_input = sigmas * noise + (1.0 - sigmas) * latents
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
# ensure the hidden state will require grad
|
||||||
|
if args.gradient_checkpointing:
|
||||||
|
noisy_model_input.requires_grad_(True)
|
||||||
|
for t in text_encoder_conds:
|
||||||
|
t.requires_grad_(True)
|
||||||
|
img_ids.requires_grad_(True)
|
||||||
|
guidance_vec.requires_grad_(True)
|
||||||
|
|
||||||
|
# Predict the noise residual
|
||||||
|
l_pooled, t5_out, txt_ids = text_encoder_conds
|
||||||
|
# print(
|
||||||
|
# f"model_input: {noisy_model_input.shape}, img_ids: {img_ids.shape}, t5_out: {t5_out.shape}, txt_ids: {txt_ids.shape}, l_pooled: {l_pooled.shape}, timesteps: {timesteps.shape}, guidance_vec: {guidance_vec.shape}"
|
||||||
|
# )
|
||||||
|
|
||||||
|
if not args.split_mode:
|
||||||
|
# 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_ids=img_ids,
|
||||||
|
txt=t5_out,
|
||||||
|
txt_ids=txt_ids,
|
||||||
|
y=l_pooled,
|
||||||
|
timesteps=timesteps / 1000,
|
||||||
|
guidance=guidance_vec,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# split forward to reduce memory usage
|
||||||
|
assert network.train_blocks == "single", "train_blocks must be single for split mode"
|
||||||
|
with accelerator.autocast():
|
||||||
|
# move flux lower to cpu, and then move flux upper to gpu
|
||||||
|
unet.to("cpu")
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
self.flux_upper.to(accelerator.device)
|
||||||
|
|
||||||
|
# upper model does not require grad
|
||||||
|
with torch.no_grad():
|
||||||
|
intermediate_img, intermediate_txt, vec, pe = self.flux_upper(
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# move flux upper back to cpu, and then move flux lower to gpu
|
||||||
|
self.flux_upper.to("cpu")
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
unet.to(accelerator.device)
|
||||||
|
|
||||||
|
# lower model requires grad
|
||||||
|
intermediate_img.requires_grad_(True)
|
||||||
|
intermediate_txt.requires_grad_(True)
|
||||||
|
vec.requires_grad_(True)
|
||||||
|
pe.requires_grad_(True)
|
||||||
|
model_pred = unet(img=intermediate_img, txt=intermediate_txt, vec=vec, pe=pe)
|
||||||
|
|
||||||
|
# unpack latents
|
||||||
|
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
|
||||||
|
|
||||||
|
if args.model_prediction_type == "raw":
|
||||||
|
# use model_pred as is
|
||||||
|
weighting = None
|
||||||
|
elif args.model_prediction_type == "additive":
|
||||||
|
# add the model_pred to the noisy_model_input
|
||||||
|
model_pred = model_pred + noisy_model_input
|
||||||
|
weighting = None
|
||||||
|
elif args.model_prediction_type == "sigma_scaled":
|
||||||
|
# apply sigma scaling
|
||||||
|
model_pred = model_pred * (-sigmas) + noisy_model_input
|
||||||
|
|
||||||
|
# these weighting schemes use a uniform timestep sampling
|
||||||
|
# and instead post-weight the loss
|
||||||
|
weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas)
|
||||||
|
|
||||||
|
# flow matching loss: this is different from SD3
|
||||||
|
target = noise - latents
|
||||||
|
|
||||||
|
return model_pred, target, timesteps, None, 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, flux="dev")
|
||||||
|
|
||||||
|
def update_metadata(self, metadata, args):
|
||||||
|
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
|
||||||
|
metadata["ss_guidance_scale"] = args.guidance_scale
|
||||||
|
metadata["ss_timestep_sampling"] = args.timestep_sampling
|
||||||
|
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
|
||||||
|
|
||||||
|
class SelectModelsTrainFlux:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"transformer": (folder_paths.get_filename_list("unet"), ),
|
||||||
|
"vae": (folder_paths.get_filename_list("vae"), ),
|
||||||
|
"clip_l": (folder_paths.get_filename_list("clip"), ),
|
||||||
|
"t5": (folder_paths.get_filename_list("clip"), ),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("TRAIN_FLUX_MODELS",)
|
||||||
|
RETURN_NAMES = ("flux_models",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "TrainFlux"
|
||||||
|
|
||||||
|
def loadmodel(self, transformer, vae, clip_l, t5):
|
||||||
|
|
||||||
|
transformer_path = folder_paths.get_full_path("unet", transformer)
|
||||||
|
vae_path = folder_paths.get_full_path("vae", vae)
|
||||||
|
clip_path = folder_paths.get_full_path("clip", clip_l)
|
||||||
|
t5_path = folder_paths.get_full_path("clip", t5)
|
||||||
|
|
||||||
|
flux_models = {
|
||||||
|
"transformer": transformer_path,
|
||||||
|
"vae": vae_path,
|
||||||
|
"clip_l": clip_path,
|
||||||
|
"t5": t5_path
|
||||||
|
}
|
||||||
|
|
||||||
|
return (flux_models,)
|
||||||
|
|
||||||
|
class TrainDatasetConfig:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"width": ("INT",{"min": 64, "default": 512}),
|
||||||
|
"height": ("INT",{"min": 64, "default": 512}),
|
||||||
|
"batch_size": ("INT",{"min": 1, "default": 2}),
|
||||||
|
"dataset_path": ("STRING",{"multiline": True, "default": ""}),
|
||||||
|
"class_tokens": ("STRING",{"multiline": True, "default": ""}),
|
||||||
|
"enable_bucket": ("BOOLEAN",{"default": True, "tooltip": "enable buckets for multi aspect ratio training"}),
|
||||||
|
"bucket_no_upscale": ("BOOLEAN",{"default": False, "tooltip": "bucket reso is defined by image size automatically"}),
|
||||||
|
"min_bucket_reso": ("INT",{"min": 64, "default": 256}),
|
||||||
|
"max_bucket_reso": ("INT",{"min": 64, "default": 1024}),
|
||||||
|
"color_aug": ("BOOLEAN",{"default": False, "tooltip": "enable weak color augmentation"}),
|
||||||
|
"flip_aug": ("BOOLEAN",{"default": False},{"tooltip": "enable horizontal flip augmentation"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("TOML_DATASET",)
|
||||||
|
RETURN_NAMES = ("dataset",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "TrainFlux"
|
||||||
|
|
||||||
|
def loadmodel(self, dataset_path, class_tokens, width, height, batch_size, enable_bucket, color_aug, flip_aug,
|
||||||
|
bucket_no_upscale, min_bucket_reso, max_bucket_reso):
|
||||||
|
import toml
|
||||||
|
|
||||||
|
dataset = {
|
||||||
|
"general": {
|
||||||
|
"shuffle_caption": False,
|
||||||
|
"caption_extension": ".txt",
|
||||||
|
},
|
||||||
|
"datasets": [
|
||||||
|
{
|
||||||
|
"resolution": (width, height),
|
||||||
|
"batch_size": batch_size,
|
||||||
|
"keep_tokens": 2,
|
||||||
|
"enable_bucket": enable_bucket,
|
||||||
|
"bucket_no_upscale": bucket_no_upscale,
|
||||||
|
"min_bucket_reso": min_bucket_reso,
|
||||||
|
"max_bucket_reso": max_bucket_reso,
|
||||||
|
"color_aug": color_aug,
|
||||||
|
"flip_aug": flip_aug,
|
||||||
|
"subsets": [
|
||||||
|
{
|
||||||
|
"image_dir": dataset_path,
|
||||||
|
"class_tokens": class_tokens
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
return (toml.dumps(dataset),)
|
||||||
|
|
||||||
|
class TrainFlux:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"flux_models": ("TRAIN_FLUX_MODELS",),
|
||||||
|
"dataset": ("TOML_DATASET",),
|
||||||
|
"output_name": ("STRING", {"default": "train_flux", "multiline": False}),
|
||||||
|
"network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}),
|
||||||
|
"learning_rate": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "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"}),
|
||||||
|
"optimizer_type": (["adamw8bit", "adafactor", "prodigy"], {"default": "adamw8bit", "tooltip": "optimizer type"}),
|
||||||
|
"max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 10000, "step": 1, "tooltip": "max number of training steps"}),
|
||||||
|
"save_every_n_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "save every n epochs"}),
|
||||||
|
"sample_every_n_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "sample every n steps"}),
|
||||||
|
"network_train_unet_only": ("BOOLEAN", {"default": True, "tooltip": "wheter to train the text encoder"}),
|
||||||
|
"text_encoder_lr": ("FLOAT", {"default": 1e-4, "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"}),
|
||||||
|
"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"}),
|
||||||
|
"weighting_scheme": (["sigma_sqrt", "logit_normal", "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"}),
|
||||||
|
"mode_scale": ("FLOAT", {"default": 1.29, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Scale of mode weighting scheme. Only effective when using the mode as the weighting_scheme"}),
|
||||||
|
"guidance_scale": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 10.0, "step": 0.01, "tooltip": "the FLUX.1 dev variant is a guidance distilled model"}),
|
||||||
|
"timestep_sampling": (["sigmoid", "uniform", "sigma"], {"tooltip": "method to sample timestep"}),
|
||||||
|
"sigmoid_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for sigmoid timestep sampling (only used when timestep-sampling is sigmoid"}),
|
||||||
|
"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)."}),
|
||||||
|
"discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}),
|
||||||
|
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
|
||||||
|
"sample_prompts": ("STRING", {"multiline": True, "default": "sample prompts", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("NETWORKTRAINER",)
|
||||||
|
RETURN_NAMES = ("network_trainer",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "TrainFlux"
|
||||||
|
|
||||||
|
def loadmodel(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, **kwargs,):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
parser = setup_parser()
|
||||||
|
args = parser.parse_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
|
||||||
|
|
||||||
|
#dataset_config = os.path.join(script_directory, "dataset_flux.toml")
|
||||||
|
output_dir = os.path.join(script_directory, "output")
|
||||||
|
if '|' in sample_prompts:
|
||||||
|
prompts = sample_prompts.split('|')
|
||||||
|
else:
|
||||||
|
prompts = [sample_prompts]
|
||||||
|
|
||||||
|
config_dict = {
|
||||||
|
"sample_prompts": prompts,
|
||||||
|
"mixed_precision": "bf16",
|
||||||
|
"num_cpu_threads_per_process": 1,
|
||||||
|
"pretrained_model_name_or_path": flux_models["transformer"],
|
||||||
|
"clip_l": flux_models["clip_l"],
|
||||||
|
"t5xxl": flux_models["t5"],
|
||||||
|
"ae": flux_models["vae"],
|
||||||
|
"save_model_as": "safetensors",
|
||||||
|
"sdpa": True,
|
||||||
|
"persistent_data_loader_workers": False,
|
||||||
|
"max_data_loader_n_workers": 0,
|
||||||
|
"seed": 42,
|
||||||
|
"gradient_checkpointing": True,
|
||||||
|
"save_precision": "bf16",
|
||||||
|
"network_module": "networks.lora_flux",
|
||||||
|
"fp8_base": True,
|
||||||
|
"dataset_config": dataset,
|
||||||
|
"output_dir": output_dir,
|
||||||
|
"output_name": output_name,
|
||||||
|
"loss_type": "l2",
|
||||||
|
"optimizer_type": optimizer_type,
|
||||||
|
}
|
||||||
|
if optimizer_type == "adafactor":
|
||||||
|
config_dict["optimizer_args"] = [
|
||||||
|
"relative_step=False",
|
||||||
|
"scale_parameter=False",
|
||||||
|
"warmup_init=False"
|
||||||
|
]
|
||||||
|
config_dict.update(kwargs)
|
||||||
|
|
||||||
|
for key, value in config_dict.items():
|
||||||
|
setattr(args, key, value)
|
||||||
|
|
||||||
|
with torch.inference_mode(False):
|
||||||
|
network_trainer = FluxNetworkTrainer()
|
||||||
|
training_loop = network_trainer.train(args)
|
||||||
|
|
||||||
|
final_output_lora_path = os.path.join(output_dir, "output", output_name)
|
||||||
|
|
||||||
|
trainer = {
|
||||||
|
"network_trainer": network_trainer,
|
||||||
|
"training_loop": training_loop,
|
||||||
|
}
|
||||||
|
return (trainer, )
|
||||||
|
|
||||||
|
class TrainLoop:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"network_trainer": ("NETWORKTRAINER",),
|
||||||
|
"epochs": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}),
|
||||||
|
"end": ("BOOLEAN", {"default": False, "tooltip": "whether to end training"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("NETWORKTRAINER", "IMAGE", "LOSSRECORDER",)
|
||||||
|
RETURN_NAMES = ("network_trainer", "validation_images", "loss_recorder")
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "TrainFlux"
|
||||||
|
|
||||||
|
def loadmodel(self, network_trainer, epochs, end):
|
||||||
|
with torch.inference_mode(False):
|
||||||
|
training_loop = network_trainer["training_loop"]
|
||||||
|
network_trainer = network_trainer["network_trainer"]
|
||||||
|
|
||||||
|
print(network_trainer.num_train_epochs)
|
||||||
|
pbar = comfy.utils.ProgressBar(epochs)
|
||||||
|
for epoch in range(epochs):
|
||||||
|
global_step, current_epoch = training_loop(
|
||||||
|
epoch=epoch,
|
||||||
|
num_train_epochs=network_trainer.num_train_epochs,
|
||||||
|
accelerator=network_trainer.accelerator,
|
||||||
|
network=network_trainer.network,
|
||||||
|
text_encoder=network_trainer.text_encoder,
|
||||||
|
unet=network_trainer.unet,
|
||||||
|
vae=network_trainer.vae,
|
||||||
|
tokenizers=network_trainer.tokenizers,
|
||||||
|
args=network_trainer.args,
|
||||||
|
train_dataloader=network_trainer.train_dataloader,
|
||||||
|
initial_step=network_trainer.initial_step,
|
||||||
|
global_step=network_trainer.global_step,
|
||||||
|
current_epoch=network_trainer.current_epoch,
|
||||||
|
metadata=network_trainer.metadata,
|
||||||
|
optimizer=network_trainer.optimizer,
|
||||||
|
lr_scheduler=network_trainer.lr_scheduler,
|
||||||
|
loss_recorder=network_trainer.loss_recorder
|
||||||
|
)
|
||||||
|
pbar.update(1)
|
||||||
|
print("GLOBAL STEP: ", global_step)
|
||||||
|
print("CURRENT EPOCH: ", current_epoch.value)
|
||||||
|
|
||||||
|
with torch.inference_mode(True):
|
||||||
|
image_tensors = flux_train_utils.sample_images(
|
||||||
|
accelerator,
|
||||||
|
network_trainer.args,
|
||||||
|
epoch,
|
||||||
|
global_step,
|
||||||
|
network_trainer.unet,
|
||||||
|
network_trainer.vae,
|
||||||
|
network_trainer.text_encoder,
|
||||||
|
network_trainer.sample_prompts_te_outputs
|
||||||
|
)
|
||||||
|
print(image_tensors.min(), image_tensors.max())
|
||||||
|
|
||||||
|
if end:
|
||||||
|
network_trainer.metadata["ss_epoch"] = str(network_trainer.num_train_epochs)
|
||||||
|
network_trainer.metadata["ss_training_finished_at"] = str(time.time())
|
||||||
|
|
||||||
|
network = accelerator.unwrap_model(network)
|
||||||
|
|
||||||
|
accelerator.end_training()
|
||||||
|
|
||||||
|
train_util.save_state_on_train_end(network_trainer.args, 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, global_step, network_trainer.num_train_epochs, force_sync_upload=True)
|
||||||
|
logger.info("model saved.")
|
||||||
|
else:
|
||||||
|
ckpt_name = train_util.get_epoch_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as, epoch + 1)
|
||||||
|
network_trainer.save_model(ckpt_name, accelerator.unwrap_model(network_trainer.network), global_step, epoch + 1)
|
||||||
|
|
||||||
|
remove_epoch_no = train_util.get_remove_epoch_no(network_trainer.args, epoch + 1)
|
||||||
|
if remove_epoch_no is not None:
|
||||||
|
remove_ckpt_name = train_util.get_epoch_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as, remove_epoch_no)
|
||||||
|
network_trainer.remove_model(remove_ckpt_name)
|
||||||
|
|
||||||
|
if network_trainer.args.save_state:
|
||||||
|
train_util.save_and_remove_state_on_epoch_end(network_trainer.args, accelerator, epoch + 1)
|
||||||
|
|
||||||
|
trainer = {
|
||||||
|
"network_trainer": network_trainer,
|
||||||
|
"training_loop": training_loop,
|
||||||
|
}
|
||||||
|
return (trainer, (0.5 * (image_tensors + 1.0)).cpu().float(), network_trainer.loss_recorder.loss_list)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"TrainFlux": TrainFlux,
|
||||||
|
"SelectModelsTrainFlux": SelectModelsTrainFlux,
|
||||||
|
"TrainDatasetConfig": TrainDatasetConfig,
|
||||||
|
"TrainLoop": TrainLoop
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"TrainFlux": "TrainFlux",
|
||||||
|
"SelectModelsTrainFlux": "SelectModelsTrainFlux",
|
||||||
|
"TrainDatasetConfig": "Train Dataset Config",
|
||||||
|
"TrainLoop": "Train Loop"
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
accelerate>=0.33.0
|
||||||
|
transformers>=4.44.0
|
||||||
|
diffusers>=0.25.0
|
||||||
|
ftfy>=6.1.1
|
||||||
|
opencv-python>=4.7.0.68
|
||||||
|
einops>=0.7.0
|
||||||
|
pytorch-lightning>=1.9.0
|
||||||
|
bitsandbytes>=0.43.3
|
||||||
|
prodigyopt>=1.0
|
||||||
|
lion-pytorch>=0.0.6
|
||||||
|
tensorboard
|
||||||
|
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
|
||||||
|
# for T5XXL tokenizer (SD3/FLUX)
|
||||||
|
sentencepiece>=0.2.0
|
||||||
+558
@@ -0,0 +1,558 @@
|
|||||||
|
# DreamBooth training
|
||||||
|
# XXX dropped option: fine_tune
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import itertools
|
||||||
|
import math
|
||||||
|
import os
|
||||||
|
from multiprocessing import Value
|
||||||
|
import toml
|
||||||
|
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from library import deepspeed_utils, strategy_base
|
||||||
|
from library.device_utils import init_ipex, clean_memory_on_device
|
||||||
|
|
||||||
|
|
||||||
|
init_ipex()
|
||||||
|
|
||||||
|
from accelerate.utils import set_seed
|
||||||
|
from diffusers import DDPMScheduler
|
||||||
|
|
||||||
|
import library.train_util as train_util
|
||||||
|
import library.config_util as config_util
|
||||||
|
from library.config_util import (
|
||||||
|
ConfigSanitizer,
|
||||||
|
BlueprintGenerator,
|
||||||
|
)
|
||||||
|
import library.custom_train_functions as custom_train_functions
|
||||||
|
from library.custom_train_functions import (
|
||||||
|
apply_snr_weight,
|
||||||
|
get_weighted_text_embeddings,
|
||||||
|
prepare_scheduler_for_custom_training,
|
||||||
|
pyramid_noise_like,
|
||||||
|
apply_noise_offset,
|
||||||
|
scale_v_prediction_loss_like_noise_prediction,
|
||||||
|
apply_debiased_estimation,
|
||||||
|
apply_masked_loss,
|
||||||
|
)
|
||||||
|
from .utils import setup_logging, add_logging_arguments
|
||||||
|
import library.strategy_sd as strategy_sd
|
||||||
|
|
||||||
|
setup_logging()
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# perlin_noise,
|
||||||
|
|
||||||
|
|
||||||
|
def train(args):
|
||||||
|
train_util.verify_training_args(args)
|
||||||
|
train_util.prepare_dataset_args(args, False)
|
||||||
|
deepspeed_utils.prepare_deepspeed_args(args)
|
||||||
|
setup_logging(args, reset=True)
|
||||||
|
|
||||||
|
cache_latents = args.cache_latents
|
||||||
|
|
||||||
|
if args.seed is not None:
|
||||||
|
set_seed(args.seed) # 乱数系列を初期化する
|
||||||
|
|
||||||
|
tokenize_strategy = strategy_sd.SdTokenizeStrategy(args.v2, args.max_token_length, args.tokenizer_cache_dir)
|
||||||
|
strategy_base.TokenizeStrategy.set_strategy(tokenize_strategy)
|
||||||
|
|
||||||
|
# prepare caching strategy: this must be set before preparing dataset. because dataset may use this strategy for initialization.
|
||||||
|
latents_caching_strategy = strategy_sd.SdSdxlLatentsCachingStrategy(
|
||||||
|
False, args.cache_latents_to_disk, args.vae_batch_size, False
|
||||||
|
)
|
||||||
|
strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy)
|
||||||
|
|
||||||
|
# データセットを準備する
|
||||||
|
if args.dataset_class is None:
|
||||||
|
blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, False, 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", "reg_data_dir"]
|
||||||
|
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:
|
||||||
|
user_config = {
|
||||||
|
"datasets": [
|
||||||
|
{"subsets": config_util.generate_dreambooth_subsets_config_by_subdirs(args.train_data_dir, args.reg_data_dir)}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
if args.no_token_padding:
|
||||||
|
train_dataset_group.disable_token_padding()
|
||||||
|
|
||||||
|
if args.debug_dataset:
|
||||||
|
train_util.debug_dataset(train_dataset_group)
|
||||||
|
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は使えません"
|
||||||
|
|
||||||
|
# acceleratorを準備する
|
||||||
|
logger.info("prepare accelerator")
|
||||||
|
|
||||||
|
if args.gradient_accumulation_steps > 1:
|
||||||
|
logger.warning(
|
||||||
|
f"gradient_accumulation_steps is {args.gradient_accumulation_steps}. accelerate does not support gradient_accumulation_steps when training multiple models (U-Net and Text Encoder), so something might be wrong"
|
||||||
|
)
|
||||||
|
logger.warning(
|
||||||
|
f"gradient_accumulation_stepsが{args.gradient_accumulation_steps}に設定されています。accelerateは複数モデル(U-NetおよびText Encoder)の学習時にgradient_accumulation_stepsをサポートしていないため結果は未知数です"
|
||||||
|
)
|
||||||
|
|
||||||
|
accelerator = train_util.prepare_accelerator(args)
|
||||||
|
|
||||||
|
# mixed precisionに対応した型を用意しておき適宜castする
|
||||||
|
weight_dtype, save_dtype = train_util.prepare_dtype(args)
|
||||||
|
vae_dtype = torch.float32 if args.no_half_vae else weight_dtype
|
||||||
|
|
||||||
|
# モデルを読み込む
|
||||||
|
text_encoder, vae, unet, load_stable_diffusion_format = train_util.load_target_model(args, weight_dtype, accelerator)
|
||||||
|
|
||||||
|
# verify load/save model formats
|
||||||
|
if load_stable_diffusion_format:
|
||||||
|
src_stable_diffusion_ckpt = args.pretrained_model_name_or_path
|
||||||
|
src_diffusers_model_path = None
|
||||||
|
else:
|
||||||
|
src_stable_diffusion_ckpt = None
|
||||||
|
src_diffusers_model_path = args.pretrained_model_name_or_path
|
||||||
|
|
||||||
|
if args.save_model_as is None:
|
||||||
|
save_stable_diffusion_format = load_stable_diffusion_format
|
||||||
|
use_safetensors = args.use_safetensors
|
||||||
|
else:
|
||||||
|
save_stable_diffusion_format = args.save_model_as.lower() == "ckpt" or args.save_model_as.lower() == "safetensors"
|
||||||
|
use_safetensors = args.use_safetensors or ("safetensors" in args.save_model_as.lower())
|
||||||
|
|
||||||
|
# モデルに xformers とか memory efficient attention を組み込む
|
||||||
|
train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers, args.sdpa)
|
||||||
|
|
||||||
|
# 学習を準備する
|
||||||
|
if cache_latents:
|
||||||
|
vae.to(accelerator.device, dtype=vae_dtype)
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
vae.eval()
|
||||||
|
|
||||||
|
train_dataset_group.new_cache_latents(vae, accelerator.is_main_process)
|
||||||
|
|
||||||
|
vae.to("cpu")
|
||||||
|
clean_memory_on_device(accelerator.device)
|
||||||
|
|
||||||
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
|
text_encoding_strategy = strategy_sd.SdTextEncodingStrategy(args.clip_skip)
|
||||||
|
strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy)
|
||||||
|
|
||||||
|
# 学習を準備する:モデルを適切な状態にする
|
||||||
|
train_text_encoder = args.stop_text_encoder_training is None or args.stop_text_encoder_training >= 0
|
||||||
|
unet.requires_grad_(True) # 念のため追加
|
||||||
|
text_encoder.requires_grad_(train_text_encoder)
|
||||||
|
if not train_text_encoder:
|
||||||
|
accelerator.print("Text Encoder is not trained.")
|
||||||
|
|
||||||
|
if args.gradient_checkpointing:
|
||||||
|
unet.enable_gradient_checkpointing()
|
||||||
|
text_encoder.gradient_checkpointing_enable()
|
||||||
|
|
||||||
|
if not cache_latents:
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
vae.eval()
|
||||||
|
vae.to(accelerator.device, dtype=weight_dtype)
|
||||||
|
|
||||||
|
# 学習に必要なクラスを準備する
|
||||||
|
accelerator.print("prepare optimizer, data loader etc.")
|
||||||
|
if train_text_encoder:
|
||||||
|
if args.learning_rate_te is None:
|
||||||
|
# wightout list, adamw8bit is crashed
|
||||||
|
trainable_params = list(itertools.chain(unet.parameters(), text_encoder.parameters()))
|
||||||
|
else:
|
||||||
|
trainable_params = [
|
||||||
|
{"params": list(unet.parameters()), "lr": args.learning_rate},
|
||||||
|
{"params": list(text_encoder.parameters()), "lr": args.learning_rate_te},
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
trainable_params = unet.parameters()
|
||||||
|
|
||||||
|
_, _, optimizer = train_util.get_optimizer(args, trainable_params)
|
||||||
|
|
||||||
|
# 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()
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
if args.stop_text_encoder_training is None:
|
||||||
|
args.stop_text_encoder_training = args.max_train_steps + 1 # do not stop until end
|
||||||
|
|
||||||
|
# lr schedulerを用意する TODO gradient_accumulation_stepsの扱いが何かおかしいかもしれない。後で確認する
|
||||||
|
lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)
|
||||||
|
|
||||||
|
# 実験的機能:勾配も含めたfp16学習を行う モデル全体をfp16にする
|
||||||
|
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.")
|
||||||
|
unet.to(weight_dtype)
|
||||||
|
text_encoder.to(weight_dtype)
|
||||||
|
|
||||||
|
# acceleratorがなんかよろしくやってくれるらしい
|
||||||
|
if args.deepspeed:
|
||||||
|
if args.train_text_encoder:
|
||||||
|
ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet, text_encoder=text_encoder)
|
||||||
|
else:
|
||||||
|
ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet)
|
||||||
|
ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||||
|
ds_model, optimizer, train_dataloader, lr_scheduler
|
||||||
|
)
|
||||||
|
training_models = [ds_model]
|
||||||
|
|
||||||
|
else:
|
||||||
|
if train_text_encoder:
|
||||||
|
unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||||
|
unet, text_encoder, optimizer, train_dataloader, lr_scheduler
|
||||||
|
)
|
||||||
|
training_models = [unet, text_encoder]
|
||||||
|
else:
|
||||||
|
unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(unet, optimizer, train_dataloader, lr_scheduler)
|
||||||
|
training_models = [unet]
|
||||||
|
|
||||||
|
if not train_text_encoder:
|
||||||
|
text_encoder.to(accelerator.device, dtype=weight_dtype) # to avoid 'cpu' vs 'cuda' error
|
||||||
|
|
||||||
|
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
|
||||||
|
if args.full_fp16:
|
||||||
|
train_util.patch_accelerator_for_fp16_training(accelerator)
|
||||||
|
|
||||||
|
# resumeする
|
||||||
|
train_util.resume_from_local_or_hf_if_specified(accelerator, args)
|
||||||
|
|
||||||
|
# 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 train images * repeats / 学習画像の数×繰り返し回数: {train_dataset_group.num_train_images}")
|
||||||
|
accelerator.print(f" num reg images / 正則化画像の数: {train_dataset_group.num_reg_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 / バッチサイズ: {args.train_batch_size}")
|
||||||
|
accelerator.print(
|
||||||
|
f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}"
|
||||||
|
)
|
||||||
|
accelerator.print(f" gradient ccumulation 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 = DDPMScheduler(
|
||||||
|
beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False
|
||||||
|
)
|
||||||
|
prepare_scheduler_for_custom_training(noise_scheduler, accelerator.device)
|
||||||
|
if args.zero_terminal_snr:
|
||||||
|
custom_train_functions.fix_noise_scheduler_betas_for_zero_terminal_snr(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(
|
||||||
|
"dreambooth" 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# For --sample_at_first
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, 0, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
loss_recorder = train_util.LossRecorder()
|
||||||
|
for epoch in range(num_train_epochs):
|
||||||
|
accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")
|
||||||
|
current_epoch.value = epoch + 1
|
||||||
|
|
||||||
|
# 指定したステップ数までText Encoderを学習する:epoch最初の状態
|
||||||
|
unet.train()
|
||||||
|
# train==True is required to enable gradient_checkpointing
|
||||||
|
if args.gradient_checkpointing or global_step < args.stop_text_encoder_training:
|
||||||
|
text_encoder.train()
|
||||||
|
|
||||||
|
for step, batch in enumerate(train_dataloader):
|
||||||
|
current_step.value = global_step
|
||||||
|
# 指定したステップ数でText Encoderの学習を止める
|
||||||
|
if global_step == args.stop_text_encoder_training:
|
||||||
|
accelerator.print(f"stop text encoder training at step {global_step}")
|
||||||
|
if not args.gradient_checkpointing:
|
||||||
|
text_encoder.train(False)
|
||||||
|
text_encoder.requires_grad_(False)
|
||||||
|
if len(training_models) == 2:
|
||||||
|
training_models = training_models[0] # remove text_encoder from training_models
|
||||||
|
|
||||||
|
with accelerator.accumulate(*training_models):
|
||||||
|
with torch.no_grad():
|
||||||
|
# latentに変換
|
||||||
|
if cache_latents:
|
||||||
|
latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)
|
||||||
|
else:
|
||||||
|
latents = vae.encode(batch["images"].to(dtype=weight_dtype)).latent_dist.sample()
|
||||||
|
latents = latents * 0.18215
|
||||||
|
b_size = latents.shape[0]
|
||||||
|
|
||||||
|
# Get the text embedding for conditioning
|
||||||
|
with torch.set_grad_enabled(global_step < args.stop_text_encoder_training):
|
||||||
|
if args.weighted_captions:
|
||||||
|
encoder_hidden_states = get_weighted_text_embeddings(
|
||||||
|
tokenize_strategy.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_ids = batch["input_ids_list"][0].to(accelerator.device)
|
||||||
|
encoder_hidden_states = text_encoding_strategy.encode_tokens(
|
||||||
|
tokenize_strategy, [text_encoder], [input_ids]
|
||||||
|
)[0]
|
||||||
|
if args.full_fp16:
|
||||||
|
encoder_hidden_states = encoder_hidden_states.to(weight_dtype)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
)
|
||||||
|
|
||||||
|
# Predict the noise residual
|
||||||
|
with accelerator.autocast():
|
||||||
|
noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample
|
||||||
|
|
||||||
|
if args.v_parameterization:
|
||||||
|
# v-parameterization training
|
||||||
|
target = noise_scheduler.get_velocity(latents, noise, timesteps)
|
||||||
|
else:
|
||||||
|
target = noise
|
||||||
|
|
||||||
|
loss = train_util.conditional_loss(
|
||||||
|
noise_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=huber_c
|
||||||
|
)
|
||||||
|
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
|
||||||
|
|
||||||
|
if args.min_snr_gamma:
|
||||||
|
loss = apply_snr_weight(loss, timesteps, noise_scheduler, args.min_snr_gamma, args.v_parameterization)
|
||||||
|
if args.scale_v_pred_loss_like_noise_pred:
|
||||||
|
loss = scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler)
|
||||||
|
if args.debiased_estimation_loss:
|
||||||
|
loss = apply_debiased_estimation(loss, timesteps, noise_scheduler)
|
||||||
|
|
||||||
|
loss = loss.mean() # 平均なのでbatch_sizeで割る必要なし
|
||||||
|
|
||||||
|
accelerator.backward(loss)
|
||||||
|
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||||
|
if train_text_encoder:
|
||||||
|
params_to_clip = itertools.chain(unet.parameters(), text_encoder.parameters())
|
||||||
|
else:
|
||||||
|
params_to_clip = unet.parameters()
|
||||||
|
accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
|
||||||
|
|
||||||
|
optimizer.step()
|
||||||
|
lr_scheduler.step()
|
||||||
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
# Checks if the accelerator has performed an optimization step behind the scenes
|
||||||
|
if accelerator.sync_gradients:
|
||||||
|
progress_bar.update(1)
|
||||||
|
global_step += 1
|
||||||
|
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, None, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
# 指定ステップごとにモデルを保存
|
||||||
|
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:
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||||
|
args,
|
||||||
|
False,
|
||||||
|
accelerator,
|
||||||
|
src_path,
|
||||||
|
save_stable_diffusion_format,
|
||||||
|
use_safetensors,
|
||||||
|
save_dtype,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
accelerator.unwrap_model(text_encoder),
|
||||||
|
accelerator.unwrap_model(unet),
|
||||||
|
vae,
|
||||||
|
)
|
||||||
|
|
||||||
|
current_loss = loss.detach().item()
|
||||||
|
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:
|
||||||
|
# checking for saving is in util
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||||
|
args,
|
||||||
|
True,
|
||||||
|
accelerator,
|
||||||
|
src_path,
|
||||||
|
save_stable_diffusion_format,
|
||||||
|
use_safetensors,
|
||||||
|
save_dtype,
|
||||||
|
epoch,
|
||||||
|
num_train_epochs,
|
||||||
|
global_step,
|
||||||
|
accelerator.unwrap_model(text_encoder),
|
||||||
|
accelerator.unwrap_model(unet),
|
||||||
|
vae,
|
||||||
|
)
|
||||||
|
|
||||||
|
train_util.sample_images(
|
||||||
|
accelerator, args, epoch + 1, global_step, accelerator.device, vae, tokenize_strategy.tokenizer, text_encoder, unet
|
||||||
|
)
|
||||||
|
|
||||||
|
is_main_process = accelerator.is_main_process
|
||||||
|
if is_main_process:
|
||||||
|
unet = accelerator.unwrap_model(unet)
|
||||||
|
text_encoder = accelerator.unwrap_model(text_encoder)
|
||||||
|
|
||||||
|
accelerator.end_training()
|
||||||
|
|
||||||
|
if is_main_process and (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:
|
||||||
|
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||||
|
train_util.save_sd_model_on_train_end(
|
||||||
|
args, src_path, save_stable_diffusion_format, use_safetensors, save_dtype, epoch, global_step, text_encoder, unet, vae
|
||||||
|
)
|
||||||
|
logger.info("model saved.")
|
||||||
|
|
||||||
|
|
||||||
|
def setup_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
add_logging_arguments(parser)
|
||||||
|
train_util.add_sd_models_arguments(parser)
|
||||||
|
train_util.add_dataset_arguments(parser, True, False, True)
|
||||||
|
train_util.add_training_arguments(parser, True)
|
||||||
|
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)
|
||||||
|
custom_train_functions.add_custom_train_arguments(parser)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--learning_rate_te",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="learning rate for text encoder, default is same as unet / Text Encoderの学習率、デフォルトはunetと同じ",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--no_token_padding",
|
||||||
|
action="store_true",
|
||||||
|
help="disable token padding (same as Diffuser's DreamBooth) / トークンのpaddingを無効にする(Diffusers版DreamBoothと同じ動作)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--stop_text_encoder_training",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="steps to stop text encoder training, -1 for no training / Text Encoderの学習を止めるステップ数、-1で最初から学習しない",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--no_half_vae",
|
||||||
|
action="store_true",
|
||||||
|
help="do not use fp16/bf16 VAE in mixed precision (use float VAE) / mixed precisionでも fp16/bf16 VAEを使わずfloat VAEを使う",
|
||||||
|
)
|
||||||
|
|
||||||
|
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)
|
||||||
+1346
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user