Compare commits
28 Commits
WIP
...
kohya-sdxl
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1d2523b978 | ||
|
|
55a5fcc7d9 | ||
|
|
e96874241d | ||
|
|
e3be1a1758 | ||
|
|
1a92e97c6d | ||
|
|
355c80df07 | ||
|
|
1487d13191 | ||
|
|
383bad958d | ||
|
|
196b693cf0 | ||
|
|
fd95e7b60c | ||
|
|
379992d89e | ||
|
|
c7054d714f | ||
|
|
67dfd9ced0 | ||
|
|
1a7e346b41 | ||
|
|
df48f0a843 | ||
|
|
fbc8a87a05 | ||
|
|
bf90740b59 | ||
|
|
ff2c9f3d04 | ||
|
|
8bd536df7e | ||
|
|
64fbd4c92a | ||
|
|
8c90fa86c6 | ||
|
|
7e4e660663 | ||
|
|
b865ac8b24 | ||
|
|
66c6f0f6f7 | ||
|
|
75ec5d9292 | ||
|
|
1a25b275c8 | ||
|
|
2bf3e529ce | ||
|
|
f53fd08690 |
4
.gitignore
vendored
4
.gitignore
vendored
@@ -170,4 +170,6 @@ cython_debug/
|
||||
!/config/examples
|
||||
!/config/_PUT_YOUR_CONFIGS_HERE).txt
|
||||
/output/*
|
||||
!/output/.gitkeep
|
||||
!/output/.gitkeep
|
||||
/extensions/*
|
||||
!/extensions/example
|
||||
59
README.md
59
README.md
@@ -29,7 +29,9 @@ cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
# or source venv/Scripts/activate on windows
|
||||
# .\venv\Scripts\activate on windows
|
||||
# windows install pytorch first with
|
||||
# pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu117
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
@@ -42,6 +44,16 @@ here so far.
|
||||
|
||||
---
|
||||
|
||||
### Batch Image Generation
|
||||
|
||||
A image generator that can take frompts from a config file or form a txt file and generate them to a
|
||||
folder. I mainly needed this for an SDXL test I am doing but added some polish to it so it can be used
|
||||
for generat batch image generation.
|
||||
It all runs off a config file, which you can find an example of in `config/examples/generate.example.yaml`.
|
||||
Mere info is in the comments in the example
|
||||
|
||||
---
|
||||
|
||||
### LoRA (lierla), LoCON (LyCORIS) extractor
|
||||
|
||||
It is based on the extractor in the [LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS) tool, but adding some QOL features
|
||||
@@ -94,6 +106,10 @@ or even -15 to 15. This will allow you to dile it in so they all have your desir
|
||||
|
||||
### LoRA Slider Trainer
|
||||
|
||||
<a target="_blank" href="https://colab.research.google.com/github/ostris/ai-toolkit/blob/main/notebooks/SliderTraining.ipynb">
|
||||
<img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/>
|
||||
</a>
|
||||
|
||||
This is how I train most of the recent sliders I have on Civitai, you can check them out in my [Civitai profile](https://civitai.com/user/Ostris/models).
|
||||
It is based off the work by [p1atdev/LECO](https://github.com/p1atdev/LECO) and [rohitgandikota/erasing](https://github.com/rohitgandikota/erasing)
|
||||
But has been heavily modified to create sliders rather than erasing concepts. I have a lot more plans on this, but it is
|
||||
@@ -116,6 +132,23 @@ I will post an better tutorial soon.
|
||||
|
||||
---
|
||||
|
||||
## Extensions!!
|
||||
|
||||
You can now make and share custom extensions. That run within this framework and have all the inbuilt tools
|
||||
available to them. I will probably use this as the primary development method going
|
||||
forward so I dont keep adding and adding more and more features to this base repo. I will likely migrate a lot
|
||||
of the existing functionality as well to make everything modular. There is an example extension in the `extensions`
|
||||
folder that shows how to make a model merger extension. All of the code is heavily documented which is hopefully
|
||||
enough to get you started. To make an extension, just copy that example and replace all the things you need to.
|
||||
|
||||
|
||||
### Model Merger - Example Extension
|
||||
It is located in the `extensions` folder. It is a fully finctional model merger that can merge as many models together
|
||||
as you want. It is a good example of how to make an extension, but is also a pretty useful feature as well since most
|
||||
mergers can only do one model at a time and this one will take as many as you want to feed it. There is an
|
||||
example config file in there, just copy that to your `config` folder and rename it to `whatever_you_want.yml`.
|
||||
and use it like any other config file.
|
||||
|
||||
## WIP Tools
|
||||
|
||||
|
||||
@@ -143,7 +176,27 @@ Just went in and out. It is much worse on smaller faces than shown here.
|
||||
|
||||
## Change Log
|
||||
|
||||
#### 2021-08-01
|
||||
#### 2023-08-05
|
||||
- Huge memory rework and slider rework. Slider training is better thant ever with no more
|
||||
ram spikes. I also made it so all 4 parts of the slider algorythm run in one batch so they share gradient
|
||||
accumulation. This makes it much faster and more stable.
|
||||
- Updated the example config to be something more practical and more updated to current methods. It is now
|
||||
a detail slide and shows how to train one without a subject. 512x512 slider training for 1.5 should work on
|
||||
6GB gpu now. Will test soon to verify.
|
||||
|
||||
|
||||
#### 2021-10-20
|
||||
- Windows support bug fixes
|
||||
- Extensions! Added functionality to make and share custom extensions for training, merging, whatever.
|
||||
check out the example in the `extensions` folder. Read more about that above.
|
||||
- Model Merging, provided via the example extension.
|
||||
|
||||
#### 2023-08-03
|
||||
Another big refactor to make SD more modular.
|
||||
|
||||
Made batch image generation script
|
||||
|
||||
#### 2023-08-01
|
||||
Major changes and update. New LoRA rescale tool, look above for details. Added better metadata so
|
||||
Automatic1111 knows what the base model is. Added some experiments and a ton of updates. This thing is still unstable
|
||||
at the moment, so hopefully there are not breaking changes.
|
||||
@@ -161,7 +214,7 @@ encoders to the model as well as a few more entirely separate diffusion networks
|
||||
training without every experimental new paper added to it. The KISS principal.
|
||||
|
||||
|
||||
#### 2021-07-30
|
||||
#### 2023-07-30
|
||||
Added "anchors" to the slider trainer. This allows you to set a prompt that will be used as a
|
||||
regularizer. You can set the network multiplier to force spread consistency at high weights
|
||||
|
||||
|
||||
60
config/examples/generate.example.yaml
Normal file
60
config/examples/generate.example.yaml
Normal file
@@ -0,0 +1,60 @@
|
||||
---
|
||||
|
||||
job: generate # tells the runner what to do
|
||||
config:
|
||||
name: "generate" # this is not really used anywhere currently but required by runner
|
||||
process:
|
||||
# process 1
|
||||
- type: to_folder # process images to a folder
|
||||
output_folder: "output/gen"
|
||||
device: cuda:0 # cpu, cuda:0, etc
|
||||
generate:
|
||||
# these are your defaults you can override most of them with flags
|
||||
sampler: "ddpm" # ignored for now, will add later though ddpm is used regardless for now
|
||||
width: 1024
|
||||
height: 1024
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
|
||||
seed: -1 # -1 is random
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
ext: ".png" # .png, .jpg, .jpeg, .webp
|
||||
|
||||
# here ate the flags you can use for prompts. Always start with
|
||||
# your prompt first then add these flags after. You can use as many
|
||||
# like
|
||||
# photo of a baseball --n painting, ugly --w 1024 --h 1024 --seed 42 --cfg 7 --steps 20
|
||||
# we will try to support all sd-scripts flags where we can
|
||||
|
||||
# FROM SD-SCRIPTS
|
||||
# --n Treat everything until the next option as a negative prompt.
|
||||
# --w Specify the width of the generated image.
|
||||
# --h Specify the height of the generated image.
|
||||
# --d Specify the seed for the generated image.
|
||||
# --l Specify the CFG scale for the generated image.
|
||||
# --s Specify the number of steps during generation.
|
||||
|
||||
# OURS and some QOL additions
|
||||
# --p2 Prompt for the second text encoder (SDXL only)
|
||||
# --n2 Negative prompt for the second text encoder (SDXL only)
|
||||
# --gr Specify the guidance rescale for the generated image (SDXL only)
|
||||
# --seed Specify the seed for the generated image same as --d
|
||||
# --cfg Specify the CFG scale for the generated image same as --l
|
||||
# --steps Specify the number of steps during generation same as --s
|
||||
|
||||
prompt_file: false # if true a txt file will be created next to images with prompt strings used
|
||||
# prompts can also be a path to a text file with one prompt per line
|
||||
# prompts: "/path/to/prompts.txt"
|
||||
prompts:
|
||||
- "photo of batman"
|
||||
- "photo of superman"
|
||||
- "photo of spiderman"
|
||||
- "photo of a superhero --n batman superman spiderman"
|
||||
|
||||
model:
|
||||
# huggingface name, relative prom project path, or absolute path to .safetensors or .ckpt
|
||||
# name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||
name_or_path: "/mnt/Models/stable-diffusion/models/stable-diffusion/Ostris/Ostris_Real_v1.safetensors"
|
||||
is_v2: false # for v2 models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
is_xl: false # for SDXL models
|
||||
dtype: bf16
|
||||
@@ -7,7 +7,7 @@ job: train
|
||||
config:
|
||||
# the name will be used to create a folder in the output folder
|
||||
# it will also replace any [name] token in the rest of this config
|
||||
name: pet_slider_v1
|
||||
name: detail_slider_v1
|
||||
# folder will be created with name above in folder below
|
||||
# it can be relative to the project root or absolute
|
||||
training_folder: "output/LoRA"
|
||||
@@ -24,7 +24,7 @@ config:
|
||||
type: "lierla"
|
||||
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
|
||||
rank: 8
|
||||
alpha: 1.0 # just leave it
|
||||
alpha: 4 # Do about half of rank
|
||||
|
||||
# training config
|
||||
train:
|
||||
@@ -33,7 +33,9 @@ config:
|
||||
# how many steps to train. More is not always better. I rarely go over 1000
|
||||
steps: 500
|
||||
# I have had good results with 4e-4 to 1e-4 at 500 steps
|
||||
lr: 1e-4
|
||||
lr: 2e-4
|
||||
# enables gradient checkpoint, saves vram, leave it on
|
||||
gradient_checkpointing: true
|
||||
# train the unet. I recommend leaving this true
|
||||
train_unet: true
|
||||
# train the text encoder. I don't recommend this unless you have a special use case
|
||||
@@ -41,6 +43,7 @@ config:
|
||||
# not the description of it (text encoder)
|
||||
train_text_encoder: false
|
||||
|
||||
|
||||
# just leave unless you know what you are doing
|
||||
# also supports "dadaptation" but set lr to 1 if you use that,
|
||||
# but it learns too fast and I don't recommend it
|
||||
@@ -51,11 +54,13 @@ config:
|
||||
# while training. Just leave it
|
||||
max_denoising_steps: 40
|
||||
# works great at 1. I do 1 even with my 4090.
|
||||
# higher may not work right with newer single batch stacking code anyway
|
||||
batch_size: 1
|
||||
# bf16 works best if your GPU supports it (modern)
|
||||
dtype: bf16 # fp32, bf16, fp16
|
||||
# if you have it, use it. It is faster and better
|
||||
xformers: true
|
||||
# torch 2.0 doesnt need xformers anymore, only use if you have lower version
|
||||
# xformers: true
|
||||
# I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX
|
||||
# although, the way we train sliders is comparative, so it probably won't work anyway
|
||||
noise_offset: 0.0
|
||||
@@ -66,11 +71,17 @@ config:
|
||||
name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||
is_v2: false # for v2 models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
# has some issues with the dual text encoder and the way we train sliders
|
||||
# it works bit weights need to probably be higher to see it.
|
||||
is_xl: false # for SDXL models
|
||||
|
||||
# saving config
|
||||
save:
|
||||
dtype: float16 # precision to save. I recommend float16
|
||||
save_every: 50 # save every this many steps
|
||||
# this will remove step counts more than this number
|
||||
# allows you to save more often in case of a crash without filling up your drive
|
||||
max_step_saves_to_keep: 2
|
||||
|
||||
# sampling config
|
||||
sample:
|
||||
@@ -88,21 +99,22 @@ config:
|
||||
# --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive
|
||||
# slide are good tests. will inherit sample.network_multiplier if not set
|
||||
# --n [string] # negative prompt, will inherit sample.neg if not set
|
||||
|
||||
# Only 75 tokens allowed currently
|
||||
prompts: # our example is an animal slider, neg: dog, pos: cat
|
||||
- "a golden retriever --m -5"
|
||||
- "a golden retriever --m -3"
|
||||
- "a golden retriever --m 3"
|
||||
- "a golden retriever --m 5"
|
||||
- "calico cat --m -5"
|
||||
- "calico cat --m -3"
|
||||
- "calico cat --m 3"
|
||||
- "calico cat --m 5"
|
||||
- "an elephant --m -5"
|
||||
- "an elephant --m -3"
|
||||
- "an elephant --m 3"
|
||||
- "an elephant --m 5"
|
||||
# I like to do a wide positive and negative spread so I can see a good range and stop
|
||||
# early if the network is braking down
|
||||
prompts:
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5"
|
||||
- "a golden retriever sitting on a leather couch, --m -5"
|
||||
- "a golden retriever sitting on a leather couch --m -3"
|
||||
- "a golden retriever sitting on a leather couch --m 3"
|
||||
- "a golden retriever sitting on a leather couch --m 5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5"
|
||||
# negative prompt used on all prompts above as default if they don't have one
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
|
||||
# seed for sampling. 42 is the answer for everything
|
||||
@@ -131,11 +143,16 @@ config:
|
||||
# resolutions to train on. [ width, height ]. This is less important for sliders
|
||||
# as we are not teaching the model anything it doesn't already know
|
||||
# but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1
|
||||
# and [ 1024, 1024 ] for sd_xl
|
||||
# you can do as many as you want here
|
||||
resolutions:
|
||||
- [ 512, 512 ]
|
||||
# - [ 512, 768 ]
|
||||
# - [ 768, 768 ]
|
||||
# slider training uses 4 combined steps for a single round. This will do it in one gradient
|
||||
# step. It is highly optimized and shouldn't take anymore vram than doing without it,
|
||||
# since we break down batches for gradient accumulation now. so just leave it on.
|
||||
batch_full_slide: true
|
||||
# These are the concepts to train on. You can do as many as you want here,
|
||||
# but they can conflict outweigh each other. Other than experimenting, I recommend
|
||||
# just doing one for good results
|
||||
@@ -146,7 +163,9 @@ config:
|
||||
# a keyword necessarily but what the model understands the concept to represent.
|
||||
# "person" will affect men, women, children, etc but will not affect cats, dogs, etc
|
||||
# it is the models base general understanding of the concept and everything it represents
|
||||
- target_class: "animal"
|
||||
# you can leave it blank to affect everything. In this example, we are adjusting
|
||||
# detail, so we will leave it blank to affect everything
|
||||
- target_class: ""
|
||||
# positive is the prompt for the positive side of the slider.
|
||||
# It is the concept that will be excited and amplified in the model when we slide the slider
|
||||
# to the positive side and forgotten / inverted when we slide
|
||||
@@ -154,33 +173,44 @@ config:
|
||||
# the prompt. You want it to be the extreme of what you want to train on. For example,
|
||||
# if you want to train on fat people, you would use "an extremely fat, morbidly obese person"
|
||||
# as the prompt. Not just "fat person"
|
||||
positive: "cat"
|
||||
# max 75 tokens for now
|
||||
positive: "high detail, 8k, intricate, detailed, high resolution, high res, high quality"
|
||||
# negative is the prompt for the negative side of the slider and works the same as positive
|
||||
# it does not necessarily work the same as a negative prompt when generating images
|
||||
negative: "dog"
|
||||
# these need to be polar opposites.
|
||||
# max 76 tokens for now
|
||||
negative: "blurry, boring, fuzzy, low detail, low resolution, low res, low quality"
|
||||
# the loss for this target is multiplied by this number.
|
||||
# if you are doing more than one target it may be good to set less important ones
|
||||
# to a lower number like 0.1 so they dont outweigh the primary target
|
||||
# to a lower number like 0.1 so they don't outweigh the primary target
|
||||
weight: 1.0
|
||||
|
||||
# anchors are prompts that wer try to hold on to while training the slider
|
||||
# you want these to generate an image very similar to the target_class
|
||||
# without directly overlapping it. For example, if you are training on a person smiling,
|
||||
# you would use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not
|
||||
anchors:
|
||||
# only positive prompt for now
|
||||
- prompt: "a woman"
|
||||
neg_prompt: "animal"
|
||||
# the multiplier applied to the LoRA when this is run.
|
||||
# higher will give it more weight but also help keep the lora from collapsing
|
||||
multiplier: 8.0
|
||||
- prompt: "a man"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
- prompt: "a person"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
|
||||
# anchors are prompts that we will try to hold on to while training the slider
|
||||
# these are NOT necessary and can prevent the slider from converging if not done right
|
||||
# leave them off if you are having issues, but they can help lock the network
|
||||
# on certain concepts to help prevent catastrophic forgetting
|
||||
# you want these to generate an image that is not your target_class, but close to it
|
||||
# is fine as long as it does not directly overlap it.
|
||||
# For example, if you are training on a person smiling,
|
||||
# you could use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not, however, the closer the concept is to the target_class
|
||||
# the less the multiplier needs to be. Keep multipliers less than 1.0 for anchors usually
|
||||
# for close concepts, you want to be closer to 0.1 or 0.2
|
||||
# these will slow down training. I am leaving them off for the demo
|
||||
|
||||
# anchors:
|
||||
# - prompt: "a woman"
|
||||
# neg_prompt: "animal"
|
||||
# # the multiplier applied to the LoRA when this is run.
|
||||
# # higher will give it more weight but also help keep the lora from collapsing
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a man"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a person"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
|
||||
# You can put any information you want here, and it will be saved in the model.
|
||||
# The below is an example, but you can put your grocery list in it if you want.
|
||||
|
||||
129
extensions/example/ExampleMergeModels.py
Normal file
129
extensions/example/ExampleMergeModels.py
Normal file
@@ -0,0 +1,129 @@
|
||||
import torch
|
||||
import gc
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from tqdm import tqdm
|
||||
|
||||
# Type check imports. Prevents circular imports
|
||||
if TYPE_CHECKING:
|
||||
from jobs import ExtensionJob
|
||||
|
||||
|
||||
# extend standard config classes to add weight
|
||||
class ModelInputConfig(ModelConfig):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.weight = kwargs.get('weight', 1.0)
|
||||
# overwrite default dtype unless user specifies otherwise
|
||||
# float 32 will give up better precision on the merging functions
|
||||
self.dtype: str = kwargs.get('dtype', 'float32')
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
# this is our main class process
|
||||
class ExampleMergeModels(BaseExtensionProcess):
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job: 'ExtensionJob',
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
# this is the setup process, do not do process intensive stuff here, just variable setup and
|
||||
# checking requirements. This is called before the run() function
|
||||
# no loading models or anything like that, it is just for setting up the process
|
||||
# all of your process intensive stuff should be done in the run() function
|
||||
# config will have everything from the process item in the config file
|
||||
|
||||
# convince methods exist on BaseProcess to get config values
|
||||
# if required is set to true and the value is not found it will throw an error
|
||||
# you can pass a default value to get_conf() as well if it was not in the config file
|
||||
# as well as a type to cast the value to
|
||||
self.save_path = self.get_conf('save_path', required=True)
|
||||
self.save_dtype = self.get_conf('save_dtype', default='float16', as_type=get_torch_dtype)
|
||||
self.device = self.get_conf('device', default='cpu', as_type=torch.device)
|
||||
|
||||
# build models to merge list
|
||||
models_to_merge = self.get_conf('models_to_merge', required=True, as_type=list)
|
||||
# build list of ModelInputConfig objects. I find it is a good idea to make a class for each config
|
||||
# this way you can add methods to it and it is easier to read and code. There are a lot of
|
||||
# inbuilt config classes located in toolkit.config_modules as well
|
||||
self.models_to_merge = [ModelInputConfig(**model) for model in models_to_merge]
|
||||
# setup is complete. Don't load anything else here, just setup variables and stuff
|
||||
|
||||
# this is the entire run process be sure to call super().run() first
|
||||
def run(self):
|
||||
# always call first
|
||||
super().run()
|
||||
print(f"Running process: {self.__class__.__name__}")
|
||||
|
||||
# let's adjust our weights first to normalize them so the total is 1.0
|
||||
total_weight = sum([model.weight for model in self.models_to_merge])
|
||||
weight_adjust = 1.0 / total_weight
|
||||
for model in self.models_to_merge:
|
||||
model.weight *= weight_adjust
|
||||
|
||||
output_model: StableDiffusion = None
|
||||
# let's do the merge, it is a good idea to use tqdm to show progress
|
||||
for model_config in tqdm(self.models_to_merge, desc="Merging models"):
|
||||
# setup model class with our helper class
|
||||
sd_model = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=model_config,
|
||||
dtype="float32"
|
||||
)
|
||||
# load the model
|
||||
sd_model.load_model()
|
||||
|
||||
# adjust the weight of the text encoder
|
||||
if isinstance(sd_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for text_encoder in sd_model.text_encoder:
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
else:
|
||||
# normal model
|
||||
for key, value in sd_model.text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
# adjust the weights of the unet
|
||||
for key, value in sd_model.unet.state_dict().items():
|
||||
value *= model_config.weight
|
||||
|
||||
if output_model is None:
|
||||
# use this one as the base
|
||||
output_model = sd_model
|
||||
else:
|
||||
# merge the models
|
||||
# text encoder
|
||||
if isinstance(output_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for i, text_encoder in enumerate(output_model.text_encoder):
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder[i].state_dict()[key]
|
||||
else:
|
||||
# normal model
|
||||
for key, value in output_model.text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder.state_dict()[key]
|
||||
# unet
|
||||
for key, value in output_model.unet.state_dict().items():
|
||||
value += sd_model.unet.state_dict()[key]
|
||||
|
||||
# remove the model to free memory
|
||||
del sd_model
|
||||
flush()
|
||||
|
||||
# merge loop is done, let's save the model
|
||||
print(f"Saving merged model to {self.save_path}")
|
||||
output_model.save(self.save_path, meta=self.meta, save_dtype=self.save_dtype)
|
||||
print(f"Saved merged model to {self.save_path}")
|
||||
# do cleanup here
|
||||
del output_model
|
||||
flush()
|
||||
25
extensions/example/__init__.py
Normal file
25
extensions/example/__init__.py
Normal file
@@ -0,0 +1,25 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# We make a subclass of Extension
|
||||
class ExampleMergeExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "example_merge_extension"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Example Merge Extension"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .ExampleMergeModels import ExampleMergeModels
|
||||
return ExampleMergeModels
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ExampleMergeExtension
|
||||
]
|
||||
48
extensions/example/config/config.example.yaml
Normal file
48
extensions/example/config/config.example.yaml
Normal file
@@ -0,0 +1,48 @@
|
||||
---
|
||||
# Always include at least one example config file to show how to use your extension.
|
||||
# use plenty of comments so users know how to use it and what everything does
|
||||
|
||||
# all extensions will use this job name
|
||||
job: extension
|
||||
config:
|
||||
name: 'my_awesome_merge'
|
||||
process:
|
||||
# Put your example processes here. This will be passed
|
||||
# to your extension process in the config argument.
|
||||
# the type MUST match your extension uid
|
||||
- type: "example_merge_extension"
|
||||
# save path for the merged model
|
||||
save_path: "output/merge/[name].safetensors"
|
||||
# save type
|
||||
dtype: fp16
|
||||
# device to run it on
|
||||
device: cuda:0
|
||||
# input models can only be SD1.x and SD2.x models for this example (currently)
|
||||
models_to_merge:
|
||||
# weights are relative, total weights will be normalized
|
||||
# for example. If you have 2 models with weight 1.0, they will
|
||||
# both be weighted 0.5. If you have 1 model with weight 1.0 and
|
||||
# another with weight 2.0, the first will be weighted 1/3 and the
|
||||
# second will be weighted 2/3
|
||||
- name_or_path: "input/model1.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model2.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model3.safetensors"
|
||||
weight: 0.3
|
||||
- name_or_path: "input/model4.safetensors"
|
||||
weight: 1.0
|
||||
|
||||
|
||||
# you can put any information you want here, and it will be saved in the model
|
||||
# the below is an example. I recommend doing trigger words at a minimum
|
||||
# in the metadata. The software will include this plus some other information
|
||||
meta:
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your model
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
@@ -0,0 +1,218 @@
|
||||
import copy
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from typing import Optional, Union, List
|
||||
from torch.utils.data import ConcatDataset, DataLoader
|
||||
from toolkit.data_loader import PairedImageDataset
|
||||
from toolkit.prompt_utils import concat_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class DatasetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.pair_folder: str = kwargs.get('pair_folder', None)
|
||||
self.network_weight: float = kwargs.get('network_weight', 1.0)
|
||||
self.target_class: str = kwargs.get('target_class', '')
|
||||
self.size: int = kwargs.get('size', 512)
|
||||
|
||||
|
||||
class ReferenceSliderConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
|
||||
self.datasets: List[DatasetConfig] = [DatasetConfig(**d) for d in kwargs.get('datasets', [])]
|
||||
|
||||
|
||||
class ImageReferenceSliderTrainerProcess(BaseSDTrainProcess):
|
||||
sd: StableDiffusion
|
||||
data_loader: DataLoader = None
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
self.prompt_txt_list = None
|
||||
self.step_num = 0
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.slider_config = ReferenceSliderConfig(**self.get_conf('slider', {}))
|
||||
|
||||
def load_datasets(self):
|
||||
if self.data_loader is None:
|
||||
print(f"Loading datasets")
|
||||
datasets = []
|
||||
for dataset in self.slider_config.datasets:
|
||||
print(f" - Dataset: {dataset.pair_folder}")
|
||||
config = {
|
||||
'path': dataset.pair_folder,
|
||||
'size': dataset.size,
|
||||
'default_prompt': dataset.target_class,
|
||||
'network_weight': dataset.network_weight,
|
||||
}
|
||||
image_dataset = PairedImageDataset(config)
|
||||
datasets.append(image_dataset)
|
||||
|
||||
concatenated_dataset = ConcatDataset(datasets)
|
||||
self.data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=self.train_config.batch_size,
|
||||
shuffle=True,
|
||||
num_workers=2
|
||||
)
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
self.load_datasets()
|
||||
|
||||
pass
|
||||
|
||||
def hook_train_loop(self, batch):
|
||||
do_mirror_loss = 'mirror' in self.slider_config.additional_losses
|
||||
|
||||
with torch.no_grad():
|
||||
imgs, prompts, base_network_weight = batch
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
imgs: torch.Tensor = imgs.to(self.device_torch, dtype=dtype)
|
||||
# split batched images in half so left is negative and right is positive
|
||||
negative_images, positive_images = torch.chunk(imgs, 2, dim=3)
|
||||
|
||||
positive_latents = self.sd.encode_images(positive_images)
|
||||
negative_latents = self.sd.encode_images(negative_images)
|
||||
|
||||
height = positive_images.shape[2]
|
||||
width = positive_images.shape[3]
|
||||
batch_size = positive_images.shape[0]
|
||||
|
||||
if self.train_config.gradient_checkpointing:
|
||||
# may get disabled elsewhere
|
||||
self.sd.unet.enable_gradient_checkpointing()
|
||||
|
||||
noise_scheduler = self.sd.noise_scheduler
|
||||
optimizer = self.optimizer
|
||||
lr_scheduler = self.lr_scheduler
|
||||
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
|
||||
timesteps = torch.randint(0, self.train_config.max_denoising_steps, (1,), device=self.device_torch)
|
||||
timesteps = timesteps.long()
|
||||
|
||||
# get noise
|
||||
noise_positive = self.sd.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
batch_size=batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
if do_mirror_loss:
|
||||
# mirror the noise
|
||||
# torch shape is [batch, channels, height, width]
|
||||
noise_negative = torch.flip(noise_positive.clone(), dims=[3])
|
||||
else:
|
||||
noise_negative = noise_positive.clone()
|
||||
|
||||
# Add noise to the latents according to the noise magnitude at each timestep
|
||||
# (this is the forward diffusion process)
|
||||
noisy_positive_latents = noise_scheduler.add_noise(positive_latents, noise_positive, timesteps)
|
||||
noisy_negative_latents = noise_scheduler.add_noise(negative_latents, noise_negative, timesteps)
|
||||
|
||||
noisy_latents = torch.cat([noisy_positive_latents, noisy_negative_latents], dim=0)
|
||||
noise = torch.cat([noise_positive, noise_negative], dim=0)
|
||||
timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
network_multiplier = [base_network_weight * 1.0, base_network_weight * -1.0]
|
||||
|
||||
flush()
|
||||
|
||||
loss_float = None
|
||||
loss_slide_float = None
|
||||
loss_mirror_float = None
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
noisy_latents.requires_grad = False
|
||||
|
||||
# if training text encoder enable grads, else do context of no grad
|
||||
with torch.set_grad_enabled(self.train_config.train_text_encoder):
|
||||
# text encoding
|
||||
embedding_list = []
|
||||
# embed the prompts
|
||||
for prompt in prompts:
|
||||
embedding = self.sd.encode_prompt(prompt).to(self.device_torch, dtype=dtype)
|
||||
embedding_list.append(embedding)
|
||||
conditional_embeds = concat_prompt_embeds(embedding_list)
|
||||
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
|
||||
self.network.multiplier = network_multiplier
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents,
|
||||
conditional_embeddings=conditional_embeds,
|
||||
timestep=timesteps,
|
||||
)
|
||||
|
||||
if self.sd.prediction_type == 'v_prediction':
|
||||
# v-parameterization training
|
||||
target = noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
|
||||
else:
|
||||
target = noise
|
||||
|
||||
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
# todo add snr gamma here
|
||||
|
||||
loss = loss.mean()
|
||||
loss_slide_float = loss.item()
|
||||
|
||||
if do_mirror_loss:
|
||||
noise_pred_pos, noise_pred_neg = torch.chunk(noise_pred, 2, dim=0)
|
||||
# mirror the negative
|
||||
noise_pred_neg = torch.flip(noise_pred_neg.clone(), dims=[3])
|
||||
loss_mirror = torch.nn.functional.mse_loss(noise_pred_pos.float(), noise_pred_neg.float(),
|
||||
reduction="none")
|
||||
loss_mirror = loss_mirror.mean([1, 2, 3])
|
||||
loss_mirror = loss_mirror.mean()
|
||||
loss_mirror_float = loss_mirror.item()
|
||||
loss += loss_mirror
|
||||
|
||||
loss_float = loss.item()
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
|
||||
flush()
|
||||
|
||||
# apply gradients
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss_float},
|
||||
)
|
||||
|
||||
if do_mirror_loss:
|
||||
loss_dict['l/s'] = loss_slide_float
|
||||
loss_dict['l/m'] = loss_mirror_float
|
||||
return loss_dict
|
||||
# end hook_train_loop
|
||||
@@ -0,0 +1,25 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# We make a subclass of Extension
|
||||
class ImageReferenceSliderTrainer(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "image_reference_slider_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Image Reference Slider Trainer"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .ImageReferenceSliderTrainerProcess import ImageReferenceSliderTrainerProcess
|
||||
return ImageReferenceSliderTrainerProcess
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ImageReferenceSliderTrainer
|
||||
]
|
||||
@@ -0,0 +1,107 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: example_name
|
||||
process:
|
||||
- type: 'image_reference_slider_trainer'
|
||||
training_folder: "/mnt/Train/out/LoRA"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "/home/jaret/Dev/.tensorboard"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 8
|
||||
linear_alpha: 8
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 5000
|
||||
lr: 1e-4
|
||||
train_unet: true
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: true
|
||||
optimizer: "adamw"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 1
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
skip_first_sample: true
|
||||
noise_offset: 0.0
|
||||
model:
|
||||
name_or_path: "/path/to/model.safetensors"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 1000 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # only affects step counts
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of a woman with red hair taking a selfie --m -3"
|
||||
- "photo of a woman with red hair taking a selfie --m -1"
|
||||
- "photo of a woman with red hair taking a selfie --m 1"
|
||||
- "photo of a woman with red hair taking a selfie --m 3"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m -3"
|
||||
- "close up photo of a man smiling at the camera, in a tank top--m -1"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m 1"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m 3"
|
||||
- "photo of a blonde woman smiling, barista --m -3"
|
||||
- "photo of a blonde woman smiling, barista --m -1"
|
||||
- "photo of a blonde woman smiling, barista --m 1"
|
||||
- "photo of a blonde woman smiling, barista --m 3"
|
||||
- "photo of a Christina Hendricks --m -1"
|
||||
- "photo of a Christina Hendricks --m -1"
|
||||
- "photo of a Christina Hendricks --m 1"
|
||||
- "photo of a Christina Hendricks --m 3"
|
||||
- "photo of a Christina Ricci --m -3"
|
||||
- "photo of a Christina Ricci --m -1"
|
||||
- "photo of a Christina Ricci --m 1"
|
||||
- "photo of a Christina Ricci --m 3"
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
|
||||
seed: 42
|
||||
walk_seed: false
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
network_multiplier: 1.0
|
||||
|
||||
logging:
|
||||
log_every: 10 # log every this many steps
|
||||
use_wandb: false # not supported yet
|
||||
verbose: false
|
||||
|
||||
slider:
|
||||
datasets:
|
||||
- pair_folder: "/path/to/folder/side/by/side/images"
|
||||
network_weight: 2.0
|
||||
target_class: "" # only used as default if caption txt are not present
|
||||
size: 512
|
||||
- pair_folder: "/path/to/folder/side/by/side/images"
|
||||
network_weight: 4.0
|
||||
target_class: "" # only used as default if caption txt are not present
|
||||
size: 512
|
||||
|
||||
|
||||
# you can put any information you want here, and it will be saved in the model
|
||||
# the below is an example. I recommend doing trigger words at a minimum
|
||||
# in the metadata. The software will include this plus some other information
|
||||
meta:
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your model
|
||||
trigger_words:
|
||||
- put
|
||||
- trigger
|
||||
- words
|
||||
- here
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
2
info.py
2
info.py
@@ -3,6 +3,6 @@ from collections import OrderedDict
|
||||
v = OrderedDict()
|
||||
v["name"] = "ai-toolkit"
|
||||
v["repo"] = "https://github.com/ostris/ai-toolkit"
|
||||
v["version"] = "0.0.2"
|
||||
v["version"] = "0.0.4"
|
||||
|
||||
software_meta = v
|
||||
|
||||
@@ -60,7 +60,11 @@ class BaseJob:
|
||||
|
||||
# check if dict key is process type
|
||||
if process['type'] in process_dict:
|
||||
ProcessClass = getattr(module, process_dict[process['type']])
|
||||
if isinstance(process_dict[process['type']], str):
|
||||
ProcessClass = getattr(module, process_dict[process['type']])
|
||||
else:
|
||||
# it is the class
|
||||
ProcessClass = process_dict[process['type']]
|
||||
self.process.append(ProcessClass(i, self, process))
|
||||
else:
|
||||
raise ValueError(f'config file is invalid. Unknown process type: {process["type"]}')
|
||||
|
||||
21
jobs/ExtensionJob.py
Normal file
21
jobs/ExtensionJob.py
Normal file
@@ -0,0 +1,21 @@
|
||||
from collections import OrderedDict
|
||||
from jobs import BaseJob
|
||||
from toolkit.extension import get_all_extensions_process_dict
|
||||
|
||||
|
||||
class ExtensionJob(BaseJob):
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
super().__init__(config)
|
||||
self.device = self.get_conf('device', 'cpu')
|
||||
self.process_dict = get_all_extensions_process_dict()
|
||||
self.load_processes(self.process_dict)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
print("")
|
||||
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
32
jobs/GenerateJob.py
Normal file
32
jobs/GenerateJob.py
Normal file
@@ -0,0 +1,32 @@
|
||||
from jobs import BaseJob
|
||||
from collections import OrderedDict
|
||||
from typing import List
|
||||
from jobs.process import GenerateProcess
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
|
||||
import sys
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
process_dict = {
|
||||
'to_folder': 'GenerateProcess',
|
||||
}
|
||||
|
||||
|
||||
class GenerateJob(BaseJob):
|
||||
process: List[GenerateProcess]
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
super().__init__(config)
|
||||
self.device = self.get_conf('device', 'cpu')
|
||||
|
||||
# loads the processes from the config
|
||||
self.load_processes(process_dict)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("")
|
||||
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
@@ -20,6 +20,8 @@ process_dict = {
|
||||
'slider_old': 'TrainSliderProcessOld',
|
||||
'lora_hack': 'TrainLoRAHack',
|
||||
'rescale_sd': 'TrainSDRescaleProcess',
|
||||
'esrgan': 'TrainESRGANProcess',
|
||||
'reference': 'TrainReferenceProcess',
|
||||
}
|
||||
|
||||
|
||||
@@ -35,18 +37,9 @@ class TrainJob(BaseJob):
|
||||
# self.mixed_precision = self.get_conf('mixed_precision', False) # fp16
|
||||
self.log_dir = self.get_conf('log_dir', None)
|
||||
|
||||
self.writer = None
|
||||
self.setup_tensorboard()
|
||||
|
||||
# loads the processes from the config
|
||||
self.load_processes(process_dict)
|
||||
|
||||
def save_training_config(self):
|
||||
timestamp = datetime.now().strftime('%Y%m%d-%H%M%S')
|
||||
os.makedirs(self.training_folder, exist_ok=True)
|
||||
save_dif = os.path.join(self.training_folder, f'run_config_{timestamp}.yaml')
|
||||
with open(save_dif, 'w') as f:
|
||||
yaml.dump(self.raw_config, f)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -55,12 +48,3 @@ class TrainJob(BaseJob):
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
|
||||
def setup_tensorboard(self):
|
||||
if self.log_dir:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
now = datetime.now()
|
||||
time_str = now.strftime('%Y%m%d-%H%M%S')
|
||||
summary_name = f"{self.name}_{time_str}"
|
||||
summary_dir = os.path.join(self.log_dir, summary_name)
|
||||
self.writer = SummaryWriter(summary_dir)
|
||||
|
||||
@@ -3,3 +3,5 @@ from .ExtractJob import ExtractJob
|
||||
from .TrainJob import TrainJob
|
||||
from .MergeJob import MergeJob
|
||||
from .ModJob import ModJob
|
||||
from .GenerateJob import GenerateJob
|
||||
from .ExtensionJob import ExtensionJob
|
||||
|
||||
20
jobs/process/BaseExtensionProcess.py
Normal file
20
jobs/process/BaseExtensionProcess.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
|
||||
|
||||
class BaseExtensionProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job,
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -1,10 +1,9 @@
|
||||
import copy
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
|
||||
|
||||
class BaseProcess:
|
||||
class BaseProcess(object):
|
||||
meta: OrderedDict
|
||||
|
||||
def __init__(
|
||||
@@ -16,6 +15,8 @@ class BaseProcess:
|
||||
self.process_id = process_id
|
||||
self.job = job
|
||||
self.config = config
|
||||
self.raw_process_config = config
|
||||
self.name = self.get_conf('name', self.job.name)
|
||||
self.meta = copy.deepcopy(self.job.meta)
|
||||
print(json.dumps(self.config, indent=4))
|
||||
|
||||
|
||||
@@ -1,34 +1,26 @@
|
||||
import glob
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from typing import Union
|
||||
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from toolkit.optimizer import get_optimizer
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
from toolkit.pipelines import CustomStableDiffusionXLPipeline, CustomStableDiffusionPipeline
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, KDPM2DiscreteScheduler, PNDMScheduler, \
|
||||
DDIMScheduler, DDPMScheduler
|
||||
from toolkit.scheduler import get_lr_scheduler
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from jobs.process import BaseTrainProcess
|
||||
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_base_model_info_to_meta
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
import gc
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from leco import train_util, model_util
|
||||
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \
|
||||
GenerateImageConfig
|
||||
|
||||
|
||||
def flush():
|
||||
@@ -36,11 +28,9 @@ def flush():
|
||||
gc.collect()
|
||||
|
||||
|
||||
UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも同じ。
|
||||
VAE_SCALE_FACTOR = 8 # 2 ** (len(vae.config.block_out_channels) - 1) = 8
|
||||
|
||||
|
||||
class BaseSDTrainProcess(BaseTrainProcess):
|
||||
sd: StableDiffusion
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, custom_pipeline=None):
|
||||
super().__init__(process_id, job, config)
|
||||
self.custom_pipeline = custom_pipeline
|
||||
@@ -48,8 +38,11 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.network_config = NetworkConfig(**self.get_conf('network', None))
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
||||
network_config = self.get_conf('network', None)
|
||||
if network_config is not None:
|
||||
self.network_config = NetworkConfig(**network_config)
|
||||
else:
|
||||
self.network_config = None
|
||||
self.train_config = TrainConfig(**self.get_conf('train', {}))
|
||||
self.model_config = ModelConfig(**self.get_conf('model', {}))
|
||||
self.save_config = SaveConfig(**self.get_conf('save', {}))
|
||||
@@ -64,177 +57,53 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
self.logging_config = LogingConfig(**self.get_conf('logging', {}))
|
||||
self.optimizer = None
|
||||
self.lr_scheduler = None
|
||||
self.sd: 'StableDiffusion' = None
|
||||
self.data_loader: Union[DataLoader, None] = None
|
||||
|
||||
# sdxl stuff
|
||||
self.logit_scale = None
|
||||
self.ckppt_info = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.train_config.dtype,
|
||||
custom_pipeline=self.custom_pipeline,
|
||||
)
|
||||
|
||||
# added later
|
||||
# to hold network if there is one
|
||||
self.network = None
|
||||
|
||||
def sample(self, step=None, is_first=False):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
os.makedirs(sample_folder, exist_ok=True)
|
||||
|
||||
if self.network is not None:
|
||||
self.network.eval()
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
original_device_dict = {
|
||||
'vae': self.sd.vae.device,
|
||||
'unet': self.sd.unet.device,
|
||||
# 'tokenizer': self.sd.tokenizer.device,
|
||||
}
|
||||
|
||||
# handle sdxl text encoder
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder, i in zip(self.sd.text_encoder, range(len(self.sd.text_encoder))):
|
||||
original_device_dict[f'text_encoder_{i}'] = encoder.device
|
||||
encoder.to(self.device_torch)
|
||||
else:
|
||||
original_device_dict['text_encoder'] = self.sd.text_encoder.device
|
||||
self.sd.text_encoder.to(self.device_torch)
|
||||
|
||||
self.sd.vae.to(self.device_torch)
|
||||
self.sd.unet.to(self.device_torch)
|
||||
# self.sd.text_encoder.to(self.device_torch)
|
||||
# self.sd.tokenizer.to(self.device_torch)
|
||||
# TODO add clip skip
|
||||
if self.sd.is_xl:
|
||||
pipeline = StableDiffusionXLPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=self.sd.noise_scheduler,
|
||||
)
|
||||
else:
|
||||
pipeline = StableDiffusionPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder,
|
||||
tokenizer=self.sd.tokenizer,
|
||||
scheduler=self.sd.noise_scheduler,
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
)
|
||||
# disable progress bar
|
||||
pipeline.set_progress_bar_config(disable=True)
|
||||
gen_img_config_list = []
|
||||
|
||||
sample_config = self.first_sample_config if is_first else self.sample_config
|
||||
|
||||
start_seed = sample_config.seed
|
||||
start_multiplier = self.network.multiplier
|
||||
current_seed = start_seed
|
||||
for i in range(len(sample_config.prompts)):
|
||||
if sample_config.walk_seed:
|
||||
current_seed = start_seed + i
|
||||
|
||||
pipeline.to(self.device_torch)
|
||||
with self.network:
|
||||
with torch.no_grad():
|
||||
if self.network is not None:
|
||||
assert self.network.is_active
|
||||
if self.logging_config.verbose:
|
||||
print("network_state", {
|
||||
'is_active': self.network.is_active,
|
||||
'multiplier': self.network.multiplier,
|
||||
})
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
|
||||
for i in tqdm(range(len(sample_config.prompts)), desc=f"Generating Samples - step: {step}",
|
||||
leave=False):
|
||||
raw_prompt = sample_config.prompts[i]
|
||||
filename = f"[time]_{step_num}_[count].png"
|
||||
|
||||
neg = sample_config.neg
|
||||
multiplier = sample_config.network_multiplier
|
||||
p_split = raw_prompt.split('--')
|
||||
prompt = p_split[0].strip()
|
||||
height = sample_config.height
|
||||
width = sample_config.width
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
|
||||
if len(p_split) > 1:
|
||||
for split in p_split:
|
||||
flag = split[:1]
|
||||
content = split[1:].strip()
|
||||
if flag == 'n':
|
||||
neg = content
|
||||
elif flag == 'm':
|
||||
# multiplier
|
||||
multiplier = float(content)
|
||||
elif flag == 'w':
|
||||
# multiplier
|
||||
width = int(content)
|
||||
elif flag == 'h':
|
||||
# multiplier
|
||||
height = int(content)
|
||||
gen_img_config_list.append(GenerateImageConfig(
|
||||
prompt=sample_config.prompts[i], # it will autoparse the prompt
|
||||
width=sample_config.width,
|
||||
height=sample_config.height,
|
||||
negative_prompt=sample_config.neg,
|
||||
seed=current_seed,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
guidance_rescale=sample_config.guidance_rescale,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
network_multiplier=sample_config.network_multiplier,
|
||||
output_path=output_path,
|
||||
))
|
||||
|
||||
height = max(64, height - height % 8) # round to divisible by 8
|
||||
width = max(64, width - width % 8) # round to divisible by 8
|
||||
|
||||
if sample_config.walk_seed:
|
||||
current_seed += i
|
||||
|
||||
if self.network is not None:
|
||||
self.network.multiplier = multiplier
|
||||
torch.manual_seed(current_seed)
|
||||
torch.cuda.manual_seed(current_seed)
|
||||
|
||||
if self.sd.is_xl:
|
||||
img = pipeline(
|
||||
prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
negative_prompt=neg,
|
||||
guidance_rescale=0.7,
|
||||
).images[0]
|
||||
else:
|
||||
img = pipeline(
|
||||
prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
negative_prompt=neg,
|
||||
).images[0]
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
seconds_since_epoch = int(time.time())
|
||||
# zero-pad 2 digits
|
||||
i_str = str(i).zfill(2)
|
||||
filename = f"{seconds_since_epoch}{step_num}_{i_str}.png"
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
img.save(output_path)
|
||||
|
||||
# clear pipeline and cache to reduce vram usage
|
||||
del pipeline
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# restore training state
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
|
||||
self.sd.vae.to(original_device_dict['vae'])
|
||||
self.sd.unet.to(original_device_dict['unet'])
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder, i in zip(self.sd.text_encoder, range(len(self.sd.text_encoder))):
|
||||
encoder.to(original_device_dict[f'text_encoder_{i}'])
|
||||
else:
|
||||
self.sd.text_encoder.to(original_device_dict['text_encoder'])
|
||||
if self.network is not None:
|
||||
self.network.train()
|
||||
self.network.multiplier = start_multiplier
|
||||
# self.sd.tokenizer.to(original_device_dict['tokenizer'])
|
||||
# send to be generated
|
||||
self.sd.generate_images(gen_img_config_list)
|
||||
|
||||
def update_training_metadata(self):
|
||||
o_dict = OrderedDict({
|
||||
@@ -328,148 +197,10 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
def hook_before_train_loop(self):
|
||||
pass
|
||||
|
||||
def get_latent_noise(
|
||||
self,
|
||||
height=None,
|
||||
width=None,
|
||||
pixel_height=None,
|
||||
pixel_width=None,
|
||||
):
|
||||
if height is None and pixel_height is None:
|
||||
raise ValueError("height or pixel_height must be specified")
|
||||
if width is None and pixel_width is None:
|
||||
raise ValueError("width or pixel_width must be specified")
|
||||
if height is None:
|
||||
height = pixel_height // VAE_SCALE_FACTOR
|
||||
if width is None:
|
||||
width = pixel_width // VAE_SCALE_FACTOR
|
||||
|
||||
noise = torch.randn(
|
||||
(
|
||||
self.train_config.batch_size,
|
||||
UNET_IN_CHANNELS,
|
||||
height,
|
||||
width,
|
||||
),
|
||||
device="cpu",
|
||||
)
|
||||
noise = apply_noise_offset(noise, self.train_config.noise_offset)
|
||||
return noise
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch=None):
|
||||
# return loss
|
||||
return 0.0
|
||||
|
||||
def get_time_ids_from_latents(self, latents):
|
||||
bs, ch, h, w = list(latents.shape)
|
||||
|
||||
height = h * VAE_SCALE_FACTOR
|
||||
width = w * VAE_SCALE_FACTOR
|
||||
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
if self.sd.is_xl:
|
||||
prompt_ids = train_util.get_add_time_ids(
|
||||
height,
|
||||
width,
|
||||
dynamic_crops=False, # look into this
|
||||
dtype=dtype,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
return train_util.concat_embeddings(
|
||||
prompt_ids, prompt_ids, bs
|
||||
)
|
||||
else:
|
||||
return None
|
||||
|
||||
def predict_noise(
|
||||
self,
|
||||
latents: torch.FloatTensor,
|
||||
text_embeddings: PromptEmbeds,
|
||||
timestep: int,
|
||||
guidance_scale=7.5,
|
||||
guidance_rescale=0, # 0.7
|
||||
add_time_ids=None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
if self.sd.is_xl:
|
||||
if add_time_ids is None:
|
||||
add_time_ids = self.get_time_ids_from_latents(latents)
|
||||
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
|
||||
latent_model_input = self.sd.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
added_cond_kwargs = {
|
||||
"text_embeds": text_embeddings.pooled_embeds,
|
||||
"time_ids": add_time_ids,
|
||||
}
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.sd.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
).sample
|
||||
|
||||
# perform guidance
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
# https://github.com/huggingface/diffusers/blob/7a91ea6c2b53f94da930a61ed571364022b21044/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py#L775
|
||||
if guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
|
||||
|
||||
else:
|
||||
# if we are doing classifier free guidance, need to double up
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
|
||||
latent_model_input = self.sd.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.sd.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
).sample
|
||||
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
return noise_pred
|
||||
|
||||
# ref: https://github.com/huggingface/diffusers/blob/0bab447670f47c28df60fbd2f6a0f833f75a16f5/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L746
|
||||
def diffuse_some_steps(
|
||||
self,
|
||||
latents: torch.FloatTensor,
|
||||
text_embeddings: PromptEmbeds,
|
||||
total_timesteps: int = 1000,
|
||||
start_timesteps=0,
|
||||
guidance_scale=1,
|
||||
add_time_ids=None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
for timestep in tqdm(self.sd.noise_scheduler.timesteps[start_timesteps:total_timesteps], leave=False):
|
||||
noise_pred = self.predict_noise(
|
||||
latents,
|
||||
text_embeddings,
|
||||
timestep,
|
||||
guidance_scale=guidance_scale,
|
||||
add_time_ids=add_time_ids,
|
||||
**kwargs,
|
||||
)
|
||||
latents = self.sd.noise_scheduler.step(noise_pred, timestep, latents).prev_sample
|
||||
|
||||
# return latents_steps
|
||||
return latents
|
||||
|
||||
def get_latest_save_path(self):
|
||||
# get latest saved step
|
||||
if os.path.exists(self.save_root):
|
||||
@@ -497,108 +228,56 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
print("load_weights not implemented for non-network models")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
# run base process run
|
||||
BaseTrainProcess.run(self)
|
||||
### HOOK ###
|
||||
self.hook_before_model_load()
|
||||
# run base sd process run
|
||||
self.sd.load_model()
|
||||
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# TODO handle other schedulers
|
||||
# sch = KDPM2DiscreteScheduler
|
||||
sch = DDPMScheduler
|
||||
# do our own scheduler
|
||||
prediction_type = "v_prediction" if self.model_config.is_v_pred else "epsilon"
|
||||
scheduler = sch(
|
||||
num_train_timesteps=1000,
|
||||
beta_start=0.00085,
|
||||
beta_end=0.0120,
|
||||
beta_schedule="scaled_linear",
|
||||
clip_sample=False,
|
||||
prediction_type=prediction_type,
|
||||
)
|
||||
if self.model_config.is_xl:
|
||||
if self.custom_pipeline is not None:
|
||||
pipln = self.custom_pipeline
|
||||
else:
|
||||
pipln = CustomStableDiffusionXLPipeline
|
||||
pipe = pipln.from_single_file(
|
||||
self.model_config.name_or_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='ddpm',
|
||||
device=self.device_torch,
|
||||
).to(self.device_torch)
|
||||
# model is loaded from BaseSDProcess
|
||||
unet = self.sd.unet
|
||||
vae = self.sd.vae
|
||||
tokenizer = self.sd.tokenizer
|
||||
text_encoder = self.sd.text_encoder
|
||||
noise_scheduler = self.sd.noise_scheduler
|
||||
|
||||
text_encoders = [pipe.text_encoder, pipe.text_encoder_2]
|
||||
tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
||||
for text_encoder in text_encoders:
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
text_encoder = text_encoders
|
||||
else:
|
||||
if self.custom_pipeline is not None:
|
||||
pipln = self.custom_pipeline
|
||||
else:
|
||||
pipln = CustomStableDiffusionPipeline
|
||||
pipe = pipln.from_single_file(
|
||||
self.model_config.name_or_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='dpm',
|
||||
device=self.device_torch,
|
||||
load_safety_checker=False,
|
||||
).to(self.device_torch)
|
||||
pipe.register_to_config(requires_safety_checker=False)
|
||||
text_encoder = pipe.text_encoder
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
tokenizer = pipe.tokenizer
|
||||
|
||||
# scheduler doesn't get set sometimes, so we set it here
|
||||
pipe.scheduler = scheduler
|
||||
|
||||
unet = pipe.unet
|
||||
noise_scheduler = pipe.scheduler
|
||||
vae = pipe.vae.to('cpu', dtype=dtype)
|
||||
vae.eval()
|
||||
vae.requires_grad_(False)
|
||||
flush()
|
||||
|
||||
self.sd = StableDiffusion(
|
||||
vae,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
unet,
|
||||
noise_scheduler,
|
||||
is_xl=self.model_config.is_xl,
|
||||
pipeline=pipe
|
||||
)
|
||||
|
||||
unet.to(self.device_torch, dtype=dtype)
|
||||
if self.train_config.xformers:
|
||||
vae.set_use_memory_efficient_attention_xformers(True)
|
||||
unet.enable_xformers_memory_efficient_attention()
|
||||
if self.train_config.gradient_checkpointing:
|
||||
unet.enable_gradient_checkpointing()
|
||||
# if isinstance(text_encoder, list):
|
||||
# for te in text_encoder:
|
||||
# te.enable_gradient_checkpointing()
|
||||
# else:
|
||||
# text_encoder.enable_gradient_checkpointing()
|
||||
|
||||
unet.to(self.device_torch, dtype=dtype)
|
||||
unet.requires_grad_(False)
|
||||
unet.eval()
|
||||
vae = vae.to(torch.device('cpu'), dtype=dtype)
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
|
||||
if self.network_config is not None:
|
||||
conv = self.network_config.conv if self.network_config.conv is not None and self.network_config.conv > 0 else None
|
||||
self.network = LoRASpecialNetwork(
|
||||
text_encoder=text_encoder,
|
||||
unet=unet,
|
||||
lora_dim=self.network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=self.network_config.alpha,
|
||||
alpha=self.network_config.linear_alpha,
|
||||
train_unet=self.train_config.train_unet,
|
||||
train_text_encoder=self.train_config.train_text_encoder,
|
||||
conv_lora_dim=conv,
|
||||
conv_alpha=self.network_config.alpha if conv is not None else None,
|
||||
conv_lora_dim=self.network_config.conv,
|
||||
conv_alpha=self.network_config.conv_alpha,
|
||||
)
|
||||
|
||||
self.network.force_to(self.device_torch, dtype=dtype)
|
||||
# give network to sd so it can use it
|
||||
self.sd.network = self.network
|
||||
|
||||
self.network.apply_to(
|
||||
text_encoder,
|
||||
@@ -615,6 +294,9 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
default_lr=self.train_config.lr
|
||||
)
|
||||
|
||||
if self.train_config.gradient_checkpointing:
|
||||
self.network.enable_gradient_checkpointing()
|
||||
|
||||
latest_save_path = self.get_latest_save_path()
|
||||
if latest_save_path is not None:
|
||||
self.print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
|
||||
@@ -641,8 +323,6 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
unet.train()
|
||||
params += unet.parameters()
|
||||
|
||||
# TODO recover save if training network. Maybe load from beginning
|
||||
|
||||
### HOOK ###
|
||||
params = self.hook_add_extra_train_params(params)
|
||||
|
||||
@@ -651,7 +331,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
optimizer_params=self.train_config.optimizer_params)
|
||||
self.optimizer = optimizer
|
||||
|
||||
lr_scheduler = train_util.get_lr_scheduler(
|
||||
lr_scheduler = get_lr_scheduler(
|
||||
self.train_config.lr_scheduler,
|
||||
optimizer,
|
||||
max_iterations=self.train_config.steps,
|
||||
@@ -682,12 +362,29 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
iterable=range(0, self.train_config.steps),
|
||||
)
|
||||
|
||||
if self.data_loader is not None:
|
||||
dataloader = self.data_loader
|
||||
dataloader_iterator = iter(dataloader)
|
||||
else:
|
||||
dataloader = None
|
||||
dataloader_iterator = None
|
||||
|
||||
# self.step_num = 0
|
||||
for step in range(self.step_num, self.train_config.steps):
|
||||
# todo handle dataloader here maybe, not sure
|
||||
if dataloader is not None:
|
||||
try:
|
||||
batch = next(dataloader_iterator)
|
||||
except StopIteration:
|
||||
# hit the end of an epoch, reset
|
||||
# todo, should we do something else here? like blow up balloons?
|
||||
dataloader_iterator = iter(dataloader)
|
||||
batch = next(dataloader_iterator)
|
||||
else:
|
||||
batch = None
|
||||
|
||||
### HOOK ###
|
||||
loss_dict = self.hook_train_loop()
|
||||
loss_dict = self.hook_train_loop(batch)
|
||||
flush()
|
||||
|
||||
if self.train_config.optimizer.lower().startswith('dadaptation') or \
|
||||
self.train_config.optimizer.lower().startswith('prodigy'):
|
||||
|
||||
@@ -1,14 +1,24 @@
|
||||
from datetime import datetime
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
import yaml
|
||||
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jobs import TrainJob, BaseJob, ExtensionJob
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class BaseTrainProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
writer: 'SummaryWriter'
|
||||
job: Union['TrainJob', 'BaseJob', 'ExtensionJob']
|
||||
progress_bar: 'tqdm' = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -18,11 +28,14 @@ class BaseTrainProcess(BaseProcess):
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.progress_bar = None
|
||||
self.writer = self.job.writer
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
||||
self.save_root = os.path.join(self.training_folder, self.job.name)
|
||||
self.writer = None
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder if hasattr(self.job, 'training_folder') else None)
|
||||
self.save_root = os.path.join(self.training_folder, self.name)
|
||||
self.step = 0
|
||||
self.first_step = 0
|
||||
self.log_dir = self.get_conf('log_dir', self.job.log_dir if hasattr(self.job, 'log_dir') else None)
|
||||
self.setup_tensorboard()
|
||||
self.save_training_config()
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -37,3 +50,19 @@ class BaseTrainProcess(BaseProcess):
|
||||
self.progress_bar.update()
|
||||
else:
|
||||
print(*args)
|
||||
|
||||
def setup_tensorboard(self):
|
||||
if self.log_dir:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
now = datetime.now()
|
||||
time_str = now.strftime('%Y%m%d-%H%M%S')
|
||||
summary_name = f"{self.name}_{time_str}"
|
||||
summary_dir = os.path.join(self.log_dir, summary_name)
|
||||
self.writer = SummaryWriter(summary_dir)
|
||||
|
||||
def save_training_config(self):
|
||||
timestamp = datetime.now().strftime('%Y%m%d-%H%M%S')
|
||||
os.makedirs(self.save_root, exist_ok=True)
|
||||
save_dif = os.path.join(self.save_root, f'process_config_{timestamp}.yaml')
|
||||
with open(save_dif, 'w') as f:
|
||||
yaml.dump(self.raw_process_config, f)
|
||||
|
||||
102
jobs/process/GenerateProcess.py
Normal file
102
jobs/process/GenerateProcess.py
Normal file
@@ -0,0 +1,102 @@
|
||||
import gc
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef, List
|
||||
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig
|
||||
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_model_hash_to_meta, \
|
||||
add_base_model_info_to_meta
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
prompts: List[str]
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.width = kwargs.get('width', 512)
|
||||
self.height = kwargs.get('height', 512)
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.prompt_2 = kwargs.get('prompt_2', None)
|
||||
self.neg_2 = kwargs.get('neg_2', None)
|
||||
self.prompts = kwargs.get('prompts', None)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.prompt_file = kwargs.get('prompt_file', False)
|
||||
if self.prompts is None:
|
||||
raise ValueError("Prompts must be set")
|
||||
if isinstance(self.prompts, str):
|
||||
if os.path.exists(self.prompts):
|
||||
with open(self.prompts, 'r', encoding='utf-8') as f:
|
||||
self.prompts = f.read().splitlines()
|
||||
self.prompts = [p.strip() for p in self.prompts if len(p.strip()) > 0]
|
||||
else:
|
||||
raise ValueError("Prompts file does not exist, put in list if you want to use a list of prompts")
|
||||
|
||||
|
||||
class GenerateProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
sd: StableDiffusion
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job,
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.output_folder = self.get_conf('output_folder', required=True)
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.model_config.dtype,
|
||||
)
|
||||
print(f"Using device {self.device}")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for prompt in self.generate_config.prompts:
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=self.generate_config.width,
|
||||
height=self.generate_config.height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs)
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
@@ -84,7 +84,8 @@ class ModRescaleLoraProcess(BaseProcess):
|
||||
if self.scale_target == 'up_down' and key.endswith('.lora_up.weight') or key.endswith('.lora_down.weight'):
|
||||
# would it be better to adjust the up weights for fp16 precision? Doing both should reduce chance of NaN
|
||||
v = v * up_down_scale
|
||||
new_state_dict[key] = v.to(get_torch_dtype(self.save_dtype))
|
||||
v = v.detach().clone().to("cpu").to(self.save_dtype)
|
||||
new_state_dict[key] = v
|
||||
|
||||
save_meta = add_model_hash_to_meta(new_state_dict, save_meta)
|
||||
save_file(new_state_dict, self.output_path, save_meta)
|
||||
|
||||
582
jobs/process/TrainESRGANProcess.py
Normal file
582
jobs/process/TrainESRGANProcess.py
Normal file
@@ -0,0 +1,582 @@
|
||||
import copy
|
||||
import glob
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
# from basicsr.archs.rrdbnet_arch import RRDBNet
|
||||
from toolkit.models.RRDB import RRDBNet as ESRGAN, esrgan_safetensors_keys
|
||||
from safetensors.torch import save_file, load_file
|
||||
from torch.utils.data import DataLoader, ConcatDataset
|
||||
import torch
|
||||
from torch import nn
|
||||
from torchvision.transforms import transforms
|
||||
|
||||
from jobs.process import BaseTrainProcess
|
||||
from toolkit.data_loader import AugmentedImageDataset
|
||||
from toolkit.esrgan_utils import convert_state_dict_to_basicsr, convert_basicsr_state_dict_to_save_format
|
||||
from toolkit.losses import ComparativeTotalVariation, get_gradient_penalty, PatternLoss
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.optimizer import get_optimizer
|
||||
from toolkit.style import get_style_model_and_losses
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from diffusers import AutoencoderKL
|
||||
from tqdm import tqdm
|
||||
import time
|
||||
import numpy as np
|
||||
from .models.vgg19_critic import Critic
|
||||
|
||||
IMAGE_TRANSFORMS = transforms.Compose(
|
||||
[
|
||||
transforms.ToTensor(),
|
||||
# transforms.Normalize([0.5], [0.5]),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class TrainESRGANProcess(BaseTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
self.data_loader = None
|
||||
self.model: ESRGAN = None
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.pretrained_path = self.get_conf('pretrained_path', 'None')
|
||||
self.datasets_objects = self.get_conf('datasets', required=True)
|
||||
self.batch_size = self.get_conf('batch_size', 1, as_type=int)
|
||||
self.resolution = self.get_conf('resolution', 256, as_type=int)
|
||||
self.learning_rate = self.get_conf('learning_rate', 1e-6, as_type=float)
|
||||
self.sample_every = self.get_conf('sample_every', None)
|
||||
self.optimizer_type = self.get_conf('optimizer', 'adam')
|
||||
self.epochs = self.get_conf('epochs', None, as_type=int)
|
||||
self.max_steps = self.get_conf('max_steps', None, as_type=int)
|
||||
self.save_every = self.get_conf('save_every', None)
|
||||
self.upscale_sample = self.get_conf('upscale_sample', 4)
|
||||
self.dtype = self.get_conf('dtype', 'float32')
|
||||
self.sample_sources = self.get_conf('sample_sources', None)
|
||||
self.log_every = self.get_conf('log_every', 100, as_type=int)
|
||||
self.style_weight = self.get_conf('style_weight', 0, as_type=float)
|
||||
self.content_weight = self.get_conf('content_weight', 0, as_type=float)
|
||||
self.mse_weight = self.get_conf('mse_weight', 1e0, as_type=float)
|
||||
self.zoom = self.get_conf('zoom', 4, as_type=int)
|
||||
self.tv_weight = self.get_conf('tv_weight', 1e0, as_type=float)
|
||||
self.critic_weight = self.get_conf('critic_weight', 1, as_type=float)
|
||||
self.pattern_weight = self.get_conf('pattern_weight', 1, as_type=float)
|
||||
self.optimizer_params = self.get_conf('optimizer_params', {})
|
||||
self.augmentations = self.get_conf('augmentations', {})
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
if self.torch_dtype == torch.bfloat16:
|
||||
self.esrgan_dtype = torch.float16
|
||||
else:
|
||||
self.esrgan_dtype = torch.float32
|
||||
self.vgg_19 = None
|
||||
self.style_weight_scalers = []
|
||||
self.content_weight_scalers = []
|
||||
|
||||
# throw error if zoom if not divisible by 2
|
||||
if self.zoom % 2 != 0:
|
||||
raise ValueError('zoom must be divisible by 2')
|
||||
|
||||
self.step_num = 0
|
||||
self.epoch_num = 0
|
||||
|
||||
self.use_critic = self.get_conf('use_critic', False, as_type=bool)
|
||||
self.critic = None
|
||||
|
||||
if self.use_critic:
|
||||
self.critic = Critic(
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
process=self,
|
||||
**self.get_conf('critic', {}) # pass any other params
|
||||
)
|
||||
|
||||
if self.sample_every is not None and self.sample_sources is None:
|
||||
raise ValueError('sample_every is specified but sample_sources is not')
|
||||
|
||||
if self.epochs is None and self.max_steps is None:
|
||||
raise ValueError('epochs or max_steps must be specified')
|
||||
|
||||
self.data_loaders = []
|
||||
# check datasets
|
||||
assert isinstance(self.datasets_objects, list)
|
||||
for dataset in self.datasets_objects:
|
||||
if 'path' not in dataset:
|
||||
raise ValueError('dataset must have a path')
|
||||
# check if is dir
|
||||
if not os.path.isdir(dataset['path']):
|
||||
raise ValueError(f"dataset path does is not a directory: {dataset['path']}")
|
||||
|
||||
# make training folder
|
||||
if not os.path.exists(self.save_root):
|
||||
os.makedirs(self.save_root, exist_ok=True)
|
||||
|
||||
self._pattern_loss = None
|
||||
|
||||
# build augmentation transforms
|
||||
aug_transforms = []
|
||||
|
||||
def update_training_metadata(self):
|
||||
self.add_meta(OrderedDict({"training_info": self.get_training_info()}))
|
||||
|
||||
def get_training_info(self):
|
||||
info = OrderedDict({
|
||||
'step': self.step_num,
|
||||
'epoch': self.epoch_num,
|
||||
})
|
||||
return info
|
||||
|
||||
def load_datasets(self):
|
||||
if self.data_loader is None:
|
||||
print(f"Loading datasets")
|
||||
datasets = []
|
||||
for dataset in self.datasets_objects:
|
||||
print(f" - Dataset: {dataset['path']}")
|
||||
ds = copy.copy(dataset)
|
||||
ds['resolution'] = self.resolution
|
||||
|
||||
if 'augmentations' not in ds:
|
||||
ds['augmentations'] = self.augmentations
|
||||
|
||||
# add the resize down augmentation
|
||||
ds['augmentations'] = [{
|
||||
'method': 'Resize',
|
||||
'params': {
|
||||
'width': int(self.resolution // self.zoom),
|
||||
'height': int(self.resolution // self.zoom),
|
||||
# downscale interpolation, string will be evaluated
|
||||
'interpolation': 'cv2.INTER_AREA'
|
||||
}
|
||||
}] + ds['augmentations']
|
||||
|
||||
image_dataset = AugmentedImageDataset(ds)
|
||||
datasets.append(image_dataset)
|
||||
|
||||
concatenated_dataset = ConcatDataset(datasets)
|
||||
self.data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=True,
|
||||
num_workers=6
|
||||
)
|
||||
|
||||
def setup_vgg19(self):
|
||||
if self.vgg_19 is None:
|
||||
self.vgg_19, self.style_losses, self.content_losses, self.vgg19_pool_4 = get_style_model_and_losses(
|
||||
single_target=True,
|
||||
device=self.device,
|
||||
output_layer_name='pool_4',
|
||||
dtype=self.torch_dtype
|
||||
)
|
||||
self.vgg_19.to(self.device, dtype=self.torch_dtype)
|
||||
self.vgg_19.requires_grad_(False)
|
||||
|
||||
# we run random noise through first to get layer scalers to normalize the loss per layer
|
||||
# bs of 2 because we run pred and target through stacked
|
||||
noise = torch.randn((2, 3, self.resolution, self.resolution), device=self.device, dtype=self.torch_dtype)
|
||||
self.vgg_19(noise)
|
||||
for style_loss in self.style_losses:
|
||||
# get a scaler to normalize to 1
|
||||
scaler = 1 / torch.mean(style_loss.loss).item()
|
||||
self.style_weight_scalers.append(scaler)
|
||||
for content_loss in self.content_losses:
|
||||
# get a scaler to normalize to 1
|
||||
scaler = 1 / torch.mean(content_loss.loss).item()
|
||||
# if is nan, set to 1
|
||||
if scaler != scaler:
|
||||
scaler = 1
|
||||
print(f"Warning: content loss scaler is nan, setting to 1")
|
||||
self.content_weight_scalers.append(scaler)
|
||||
|
||||
self.print(f"Style weight scalers: {self.style_weight_scalers}")
|
||||
self.print(f"Content weight scalers: {self.content_weight_scalers}")
|
||||
|
||||
def get_style_loss(self):
|
||||
if self.style_weight > 0:
|
||||
# scale all losses with loss scalers
|
||||
loss = torch.sum(
|
||||
torch.stack([loss.loss * scaler for loss, scaler in zip(self.style_losses, self.style_weight_scalers)]))
|
||||
return loss
|
||||
else:
|
||||
return torch.tensor(0.0, device=self.device)
|
||||
|
||||
def get_content_loss(self):
|
||||
if self.content_weight > 0:
|
||||
# scale all losses with loss scalers
|
||||
loss = torch.sum(torch.stack(
|
||||
[loss.loss * scaler for loss, scaler in zip(self.content_losses, self.content_weight_scalers)]))
|
||||
return loss
|
||||
else:
|
||||
return torch.tensor(0.0, device=self.device)
|
||||
|
||||
def get_mse_loss(self, pred, target):
|
||||
if self.mse_weight > 0:
|
||||
loss_fn = nn.MSELoss()
|
||||
loss = loss_fn(pred, target)
|
||||
return loss
|
||||
else:
|
||||
return torch.tensor(0.0, device=self.device)
|
||||
|
||||
def get_tv_loss(self, pred, target):
|
||||
if self.tv_weight > 0:
|
||||
get_tv_loss = ComparativeTotalVariation()
|
||||
loss = get_tv_loss(pred, target)
|
||||
return loss
|
||||
else:
|
||||
return torch.tensor(0.0, device=self.device)
|
||||
|
||||
def get_pattern_loss(self, pred, target):
|
||||
if self._pattern_loss is None:
|
||||
self._pattern_loss = PatternLoss(
|
||||
pattern_size=self.zoom,
|
||||
dtype=self.torch_dtype
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
loss = torch.mean(self._pattern_loss(pred, target))
|
||||
return loss
|
||||
|
||||
def save(self, step=None):
|
||||
if not os.path.exists(self.save_root):
|
||||
os.makedirs(self.save_root, exist_ok=True)
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zeropad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
|
||||
self.update_training_metadata()
|
||||
# filename = f'{self.job.name}{step_num}.safetensors'
|
||||
filename = f'{self.job.name}{step_num}.pth'
|
||||
# prepare meta
|
||||
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
|
||||
|
||||
# state_dict = self.model.state_dict()
|
||||
|
||||
# state has the original state dict keys so we can save what we started from
|
||||
save_state_dict = self.model.state_dict()
|
||||
|
||||
for key in list(save_state_dict.keys()):
|
||||
v = save_state_dict[key]
|
||||
v = v.detach().clone().to("cpu").to(torch.float32)
|
||||
save_state_dict[key] = v
|
||||
|
||||
# most things wont use safetensors, save as torch
|
||||
# save_file(save_state_dict, os.path.join(self.save_root, filename), save_meta)
|
||||
torch.save(save_state_dict, os.path.join(self.save_root, filename))
|
||||
|
||||
self.print(f"Saved to {os.path.join(self.save_root, filename)}")
|
||||
|
||||
if self.use_critic:
|
||||
self.critic.save(step)
|
||||
|
||||
def sample(self, step=None):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
os.makedirs(sample_folder, exist_ok=True)
|
||||
|
||||
self.model.eval()
|
||||
|
||||
with torch.no_grad():
|
||||
for i, img_url in enumerate(self.sample_sources):
|
||||
img = exif_transpose(Image.open(img_url))
|
||||
img = img.convert('RGB')
|
||||
# crop if not square
|
||||
if img.width != img.height:
|
||||
min_dim = min(img.width, img.height)
|
||||
img = img.crop((0, 0, min_dim, min_dim))
|
||||
# resize
|
||||
img = img.resize((self.resolution * self.zoom, self.resolution * self.zoom), resample=Image.BICUBIC)
|
||||
|
||||
target_image = img
|
||||
# downscale the image input
|
||||
img = img.resize((self.resolution, self.resolution), resample=Image.BICUBIC)
|
||||
|
||||
# downscale the image input
|
||||
|
||||
img = IMAGE_TRANSFORMS(img).unsqueeze(0).to(self.device, dtype=self.esrgan_dtype)
|
||||
img = img
|
||||
output = self.model(img)
|
||||
# output = (output / 2 + 0.5).clamp(0, 1)
|
||||
output = output.clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
output = output.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
|
||||
|
||||
# convert to pillow image
|
||||
output = Image.fromarray((output * 255).astype(np.uint8))
|
||||
|
||||
# upscale to size * self.upscale_sample while maintaining pixels
|
||||
output = output.resize(
|
||||
(self.resolution * self.upscale_sample, self.resolution * self.upscale_sample),
|
||||
resample=Image.NEAREST
|
||||
)
|
||||
|
||||
width, height = output.size
|
||||
|
||||
# stack input image and decoded image
|
||||
target_image = target_image.resize((width, height))
|
||||
output = output.resize((width, height))
|
||||
|
||||
output_img = Image.new('RGB', (width * 2, height))
|
||||
output_img.paste(target_image, (0, 0))
|
||||
output_img.paste(output, (width, 0))
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
seconds_since_epoch = int(time.time())
|
||||
# zero-pad 2 digits
|
||||
i_str = str(i).zfill(2)
|
||||
filename = f"{seconds_since_epoch}{step_num}_{i_str}.png"
|
||||
output_img.save(os.path.join(sample_folder, filename))
|
||||
|
||||
self.model.train()
|
||||
|
||||
def load_model(self):
|
||||
state_dict = None
|
||||
path_to_load = self.pretrained_path
|
||||
# see if we have a checkpoint in out output to resume from
|
||||
self.print(f"Looking for latest checkpoint in {self.save_root}")
|
||||
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*.safetensors"))
|
||||
files += glob.glob(os.path.join(self.save_root, f"{self.job.name}*.pth"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
path_to_load = latest_file
|
||||
# todo update step and epoch count
|
||||
elif self.pretrained_path is None:
|
||||
self.print(f" - No checkpoint found, starting from scratch")
|
||||
else:
|
||||
self.print(f" - No checkpoint found, loading pretrained model")
|
||||
self.print(f" - path: {path_to_load}")
|
||||
|
||||
if path_to_load is not None:
|
||||
self.print(f" - Loading pretrained checkpoint: {path_to_load}")
|
||||
# if ends with pth then assume pytorch checkpoint
|
||||
if path_to_load.endswith('.pth') or path_to_load.endswith('.pt'):
|
||||
state_dict = torch.load(path_to_load, map_location=self.device)
|
||||
elif path_to_load.endswith('.safetensors'):
|
||||
state_dict_raw = load_file(path_to_load)
|
||||
# make ordered dict as most things need it
|
||||
state_dict = OrderedDict()
|
||||
for key in esrgan_safetensors_keys:
|
||||
state_dict[key] = state_dict_raw[key]
|
||||
else:
|
||||
raise Exception(f"Unknown file extension for checkpoint: {path_to_load}")
|
||||
|
||||
# todo determine architecture from checkpoint
|
||||
self.model = ESRGAN(
|
||||
state_dict
|
||||
).to(self.device, dtype=self.esrgan_dtype)
|
||||
|
||||
# set the model to training mode
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
self.load_datasets()
|
||||
|
||||
max_step_epochs = self.max_steps // len(self.data_loader)
|
||||
num_epochs = self.epochs
|
||||
if num_epochs is None or num_epochs > max_step_epochs:
|
||||
num_epochs = max_step_epochs
|
||||
|
||||
max_epoch_steps = len(self.data_loader) * num_epochs
|
||||
num_steps = self.max_steps
|
||||
if num_steps is None or num_steps > max_epoch_steps:
|
||||
num_steps = max_epoch_steps
|
||||
self.max_steps = num_steps
|
||||
self.epochs = num_epochs
|
||||
start_step = self.step_num
|
||||
self.first_step = start_step
|
||||
|
||||
self.print(f"Training ESRGAN model:")
|
||||
self.print(f" - Training folder: {self.training_folder}")
|
||||
self.print(f" - Batch size: {self.batch_size}")
|
||||
self.print(f" - Learning rate: {self.learning_rate}")
|
||||
self.print(f" - Epochs: {num_epochs}")
|
||||
self.print(f" - Max steps: {self.max_steps}")
|
||||
|
||||
# load model
|
||||
self.load_model()
|
||||
|
||||
params = self.model.parameters()
|
||||
|
||||
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
|
||||
self.setup_vgg19()
|
||||
self.vgg_19.requires_grad_(False)
|
||||
self.vgg_19.eval()
|
||||
if self.use_critic:
|
||||
self.critic.setup()
|
||||
|
||||
optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
|
||||
# setup scheduler
|
||||
# todo allow other schedulers
|
||||
scheduler = torch.optim.lr_scheduler.ConstantLR(
|
||||
optimizer,
|
||||
total_iters=num_steps,
|
||||
factor=1,
|
||||
verbose=False
|
||||
)
|
||||
|
||||
# setup tqdm progress bar
|
||||
self.progress_bar = tqdm(
|
||||
total=num_steps,
|
||||
desc='Training ESRGAN',
|
||||
leave=True
|
||||
)
|
||||
|
||||
blank_losses = OrderedDict({
|
||||
"total": [],
|
||||
"style": [],
|
||||
"content": [],
|
||||
"mse": [],
|
||||
"kl": [],
|
||||
"tv": [],
|
||||
"ptn": [],
|
||||
"crD": [],
|
||||
"crG": [],
|
||||
})
|
||||
epoch_losses = copy.deepcopy(blank_losses)
|
||||
log_losses = copy.deepcopy(blank_losses)
|
||||
print("Generating baseline samples")
|
||||
self.sample(step=0)
|
||||
# range start at self.epoch_num go to self.epochs
|
||||
for epoch in range(self.epoch_num, self.epochs, 1):
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
for targets, inputs in self.data_loader:
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
with torch.no_grad():
|
||||
targets = targets.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1)
|
||||
inputs = inputs.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1)
|
||||
|
||||
pred = self.model(inputs)
|
||||
|
||||
pred = pred.to(self.device, dtype=self.torch_dtype).clamp(0, 1)
|
||||
targets = targets.to(self.device, dtype=self.torch_dtype).clamp(0, 1)
|
||||
|
||||
# Run through VGG19
|
||||
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
|
||||
stacked = torch.cat([pred, targets], dim=0)
|
||||
# stacked = (stacked / 2 + 0.5).clamp(0, 1)
|
||||
stacked = stacked.clamp(0, 1)
|
||||
self.vgg_19(stacked)
|
||||
|
||||
if self.use_critic:
|
||||
critic_d_loss = self.critic.step(self.vgg19_pool_4.tensor.detach())
|
||||
else:
|
||||
critic_d_loss = 0.0
|
||||
|
||||
style_loss = self.get_style_loss() * self.style_weight
|
||||
content_loss = self.get_content_loss() * self.content_weight
|
||||
mse_loss = self.get_mse_loss(pred, targets) * self.mse_weight
|
||||
tv_loss = self.get_tv_loss(pred, targets) * self.tv_weight
|
||||
pattern_loss = self.get_pattern_loss(pred, targets) * self.pattern_weight
|
||||
if self.use_critic:
|
||||
critic_gen_loss = self.critic.get_critic_loss(self.vgg19_pool_4.tensor) * self.critic_weight
|
||||
else:
|
||||
critic_gen_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
|
||||
|
||||
loss = style_loss + content_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss
|
||||
|
||||
# Backward pass and optimization
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
scheduler.step()
|
||||
|
||||
# update progress bar
|
||||
loss_value = loss.item()
|
||||
# get exponent like 3.54e-4
|
||||
loss_string = f"loss: {loss_value:.2e}"
|
||||
if self.content_weight > 0:
|
||||
loss_string += f" cnt: {content_loss.item():.2e}"
|
||||
if self.style_weight > 0:
|
||||
loss_string += f" sty: {style_loss.item():.2e}"
|
||||
if self.mse_weight > 0:
|
||||
loss_string += f" mse: {mse_loss.item():.2e}"
|
||||
if self.tv_weight > 0:
|
||||
loss_string += f" tv: {tv_loss.item():.2e}"
|
||||
if self.pattern_weight > 0:
|
||||
loss_string += f" ptn: {pattern_loss.item():.2e}"
|
||||
if self.use_critic and self.critic_weight > 0:
|
||||
loss_string += f" crG: {critic_gen_loss.item():.2e}"
|
||||
if self.use_critic:
|
||||
loss_string += f" crD: {critic_d_loss:.2e}"
|
||||
|
||||
if self.optimizer_type.startswith('dadaptation') or self.optimizer_type.startswith('prodigy'):
|
||||
learning_rate = (
|
||||
optimizer.param_groups[0]["d"] *
|
||||
optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
learning_rate = optimizer.param_groups[0]['lr']
|
||||
|
||||
lr_critic_string = ''
|
||||
if self.use_critic:
|
||||
lr_critic = self.critic.get_lr()
|
||||
lr_critic_string = f" lrC: {lr_critic:.1e}"
|
||||
|
||||
self.progress_bar.set_postfix_str(f"lr: {learning_rate:.1e}{lr_critic_string} {loss_string}")
|
||||
self.progress_bar.set_description(f"E: {epoch}")
|
||||
self.progress_bar.update(1)
|
||||
|
||||
epoch_losses["total"].append(loss_value)
|
||||
epoch_losses["style"].append(style_loss.item())
|
||||
epoch_losses["content"].append(content_loss.item())
|
||||
epoch_losses["mse"].append(mse_loss.item())
|
||||
epoch_losses["tv"].append(tv_loss.item())
|
||||
epoch_losses["ptn"].append(pattern_loss.item())
|
||||
epoch_losses["crG"].append(critic_gen_loss.item())
|
||||
epoch_losses["crD"].append(critic_d_loss)
|
||||
|
||||
log_losses["total"].append(loss_value)
|
||||
log_losses["style"].append(style_loss.item())
|
||||
log_losses["content"].append(content_loss.item())
|
||||
log_losses["mse"].append(mse_loss.item())
|
||||
log_losses["tv"].append(tv_loss.item())
|
||||
log_losses["ptn"].append(pattern_loss.item())
|
||||
log_losses["crG"].append(critic_gen_loss.item())
|
||||
log_losses["crD"].append(critic_d_loss)
|
||||
|
||||
# don't do on first step
|
||||
if self.step_num != start_step:
|
||||
if self.sample_every and self.step_num % self.sample_every == 0:
|
||||
# print above the progress bar
|
||||
self.print(f"Sampling at step {self.step_num}")
|
||||
self.sample(self.step_num)
|
||||
|
||||
if self.save_every and self.step_num % self.save_every == 0:
|
||||
# print above the progress bar
|
||||
self.print(f"Saving at step {self.step_num}")
|
||||
self.save(self.step_num)
|
||||
|
||||
if self.log_every and self.step_num % self.log_every == 0:
|
||||
# log to tensorboard
|
||||
if self.writer is not None:
|
||||
# get avg loss
|
||||
for key in log_losses:
|
||||
log_losses[key] = sum(log_losses[key]) / (len(log_losses[key]) + 1e-6)
|
||||
# if log_losses[key] > 0:
|
||||
self.writer.add_scalar(f"loss/{key}", log_losses[key], self.step_num)
|
||||
# reset log losses
|
||||
log_losses = copy.deepcopy(blank_losses)
|
||||
|
||||
self.step_num += 1
|
||||
# end epoch
|
||||
if self.writer is not None:
|
||||
eps = 1e-6
|
||||
# get avg loss
|
||||
for key in epoch_losses:
|
||||
epoch_losses[key] = sum(log_losses[key]) / (len(log_losses[key]) + eps)
|
||||
if epoch_losses[key] > 0:
|
||||
self.writer.add_scalar(f"epoch loss/{key}", epoch_losses[key], epoch)
|
||||
# reset epoch losses
|
||||
epoch_losses = copy.deepcopy(blank_losses)
|
||||
|
||||
self.save()
|
||||
@@ -68,7 +68,7 @@ class TrainLoRAHack(BaseSDTrainProcess):
|
||||
|
||||
return loss_dict
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
if self.hack_config.type == 'suppression':
|
||||
return self.supress_loop()
|
||||
else:
|
||||
|
||||
@@ -1,24 +1,14 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import glob
|
||||
import os
|
||||
from typing import Optional
|
||||
from collections import OrderedDict
|
||||
import random
|
||||
from typing import Optional, List
|
||||
|
||||
import numpy as np
|
||||
from safetensors.torch import load_file, save_file
|
||||
from safetensors.torch import save_file, load_file
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.layers import ReductionKernel
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||
from toolkit.train_pipelines import TransferStableDiffusionXLPipeline
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
@@ -40,12 +30,10 @@ class RescaleConfig:
|
||||
):
|
||||
self.from_resolution = kwargs.get('from_resolution', 512)
|
||||
self.scale = kwargs.get('scale', 0.5)
|
||||
self.prompt_file = kwargs.get('prompt_file', None)
|
||||
self.prompt_tensors = kwargs.get('prompt_tensors', None)
|
||||
self.latent_tensor_dir = kwargs.get('latent_tensor_dir', None)
|
||||
self.num_latent_tensors = kwargs.get('num_latent_tensors', 1000)
|
||||
self.to_resolution = kwargs.get('to_resolution', int(self.from_resolution * self.scale))
|
||||
|
||||
if self.prompt_file is None:
|
||||
raise ValueError("prompt_file is required")
|
||||
self.prompt_dropout = kwargs.get('prompt_dropout', 0.1)
|
||||
|
||||
|
||||
class PromptEmbedsCache:
|
||||
@@ -64,12 +52,11 @@ class PromptEmbedsCache:
|
||||
class TrainSDRescaleProcess(BaseSDTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
# pass our custom pipeline to super so it sets it up
|
||||
super().__init__(process_id, job, config, custom_pipeline=TransferStableDiffusionXLPipeline)
|
||||
super().__init__(process_id, job, config)
|
||||
self.step_num = 0
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.prompt_cache = PromptEmbedsCache()
|
||||
self.rescale_config = RescaleConfig(**self.get_conf('rescale', required=True))
|
||||
self.reduce_size_fn = ReductionKernel(
|
||||
in_channels=4,
|
||||
@@ -77,190 +64,213 @@ class TrainSDRescaleProcess(BaseSDTrainProcess):
|
||||
dtype=get_torch_dtype(self.train_config.dtype),
|
||||
device=self.device_torch,
|
||||
)
|
||||
self.prompt_txt_list = []
|
||||
|
||||
self.latent_paths: List[str] = []
|
||||
self.empty_embedding: PromptEmbeds = None
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def get_latent_tensors(self):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
num_to_generate = 0
|
||||
# check if dir exists
|
||||
if not os.path.exists(self.rescale_config.latent_tensor_dir):
|
||||
os.makedirs(self.rescale_config.latent_tensor_dir)
|
||||
num_to_generate = self.rescale_config.num_latent_tensors
|
||||
else:
|
||||
# find existing
|
||||
current_tensor_list = glob.glob(os.path.join(self.rescale_config.latent_tensor_dir, "*.safetensors"))
|
||||
num_to_generate = self.rescale_config.num_latent_tensors - len(current_tensor_list)
|
||||
self.latent_paths = current_tensor_list
|
||||
|
||||
if num_to_generate > 0:
|
||||
print(f"Generating {num_to_generate}/{self.rescale_config.num_latent_tensors} latent tensors")
|
||||
|
||||
# unload other model
|
||||
self.sd.unet.to('cpu')
|
||||
|
||||
# load aux network
|
||||
self.sd_parent = StableDiffusion(
|
||||
self.device_torch,
|
||||
model_config=self.model_config,
|
||||
dtype=self.train_config.dtype,
|
||||
)
|
||||
self.sd_parent.load_model()
|
||||
self.sd_parent.unet.to(self.device_torch, dtype=dtype)
|
||||
# we dont need text encoder for this
|
||||
|
||||
del self.sd_parent.text_encoder
|
||||
del self.sd_parent.tokenizer
|
||||
|
||||
self.sd_parent.unet.eval()
|
||||
self.sd_parent.unet.requires_grad_(False)
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
self.empty_embedding, # unconditional (negative prompt)
|
||||
self.empty_embedding, # conditional (positive prompt)
|
||||
self.train_config.batch_size,
|
||||
)
|
||||
torch.set_default_device(self.device_torch)
|
||||
|
||||
for i in tqdm(range(num_to_generate)):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
# get a random seed
|
||||
seed = torch.randint(0, 2 ** 32, (1,)).item()
|
||||
# zero pad seed string to max length
|
||||
seed_string = str(seed).zfill(10)
|
||||
# set seed
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
# # ger a random number of steps
|
||||
timesteps_to = self.train_config.max_denoising_steps
|
||||
|
||||
# set the scheduler to the number of steps
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
timesteps_to, device=self.device_torch
|
||||
)
|
||||
|
||||
noise = self.sd.get_latent_noise(
|
||||
pixel_height=self.rescale_config.from_resolution,
|
||||
pixel_width=self.rescale_config.from_resolution,
|
||||
batch_size=self.train_config.batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get latents
|
||||
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
||||
latents = latents.to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get random guidance scale from 1.0 to 10.0 (CFG)
|
||||
guidance_scale = torch.rand(1).item() * 9.0 + 1.0
|
||||
|
||||
# do a timestep of 1
|
||||
timestep = 1
|
||||
|
||||
noise_pred_target = self.sd_parent.predict_noise(
|
||||
latents,
|
||||
text_embeddings=text_embeddings,
|
||||
timestep=timestep,
|
||||
guidance_scale=guidance_scale
|
||||
)
|
||||
|
||||
# build state dict
|
||||
state_dict = OrderedDict()
|
||||
state_dict['noise_pred_target'] = noise_pred_target.to('cpu', dtype=torch.float16)
|
||||
state_dict['latents'] = latents.to('cpu', dtype=torch.float16)
|
||||
state_dict['guidance_scale'] = torch.tensor(guidance_scale).to('cpu', dtype=torch.float16)
|
||||
state_dict['timestep'] = torch.tensor(timestep).to('cpu', dtype=torch.float16)
|
||||
state_dict['timesteps_to'] = torch.tensor(timesteps_to).to('cpu', dtype=torch.float16)
|
||||
state_dict['seed'] = torch.tensor(seed).to('cpu', dtype=torch.float32) # must be float 32 to prevent overflow
|
||||
|
||||
file_name = f"{seed_string}_{i}.safetensors"
|
||||
file_path = os.path.join(self.rescale_config.latent_tensor_dir, file_name)
|
||||
save_file(state_dict, file_path)
|
||||
self.latent_paths.append(file_path)
|
||||
|
||||
print("Removing parent model")
|
||||
# delete parent
|
||||
del self.sd_parent
|
||||
flush()
|
||||
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
self.sd.unet.to(self.device_torch, dtype=dtype)
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.print(f"Loading prompt file from {self.rescale_config.prompt_file}")
|
||||
# encode our empty prompt
|
||||
self.empty_embedding = self.sd.encode_prompt("")
|
||||
self.empty_embedding = self.empty_embedding.to(self.device_torch,
|
||||
dtype=get_torch_dtype(self.train_config.dtype))
|
||||
|
||||
# read line by line from file
|
||||
with open(self.rescale_config.prompt_file, 'r') as f:
|
||||
self.prompt_txt_list = f.readlines()
|
||||
# clean empty lines
|
||||
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
|
||||
|
||||
self.print(f"Loaded {len(self.prompt_txt_list)} prompts. Encoding them..")
|
||||
|
||||
cache = PromptEmbedsCache()
|
||||
|
||||
# get encoded latents for our prompts
|
||||
with torch.no_grad():
|
||||
if self.rescale_config.prompt_tensors is not None:
|
||||
# check to see if it exists
|
||||
if os.path.exists(self.rescale_config.prompt_tensors):
|
||||
# load it.
|
||||
self.print(f"Loading prompt tensors from {self.rescale_config.prompt_tensors}")
|
||||
prompt_tensors = load_file(self.rescale_config.prompt_tensors, device='cpu')
|
||||
# add them to the cache
|
||||
for prompt_txt, prompt_tensor in prompt_tensors.items():
|
||||
if prompt_txt.startswith("te:"):
|
||||
prompt = prompt_txt[3:]
|
||||
# text_embeds
|
||||
text_embeds = prompt_tensor
|
||||
pooled_embeds = None
|
||||
# find pool embeds
|
||||
if f"pe:{prompt}" in prompt_tensors:
|
||||
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
|
||||
|
||||
# make it
|
||||
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
|
||||
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
|
||||
|
||||
if len(cache.prompts) == 0:
|
||||
print("Prompt tensors not found. Encoding prompts..")
|
||||
neutral = ""
|
||||
# encode neutral
|
||||
cache[neutral] = self.sd.encode_prompt(neutral)
|
||||
for prompt in tqdm(self.prompt_txt_list, desc="Encoding prompts", leave=False):
|
||||
# build the cache
|
||||
if cache[prompt] is None:
|
||||
cache[prompt] = self.sd.encode_prompt(prompt).to(device="cpu", dtype=torch.float32)
|
||||
|
||||
if self.rescale_config.prompt_tensors:
|
||||
print(f"Saving prompt tensors to {self.rescale_config.prompt_tensors}")
|
||||
state_dict = {}
|
||||
for prompt_txt, prompt_embeds in cache.prompts.items():
|
||||
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to("cpu",
|
||||
dtype=get_torch_dtype('fp16'))
|
||||
if prompt_embeds.pooled_embeds is not None:
|
||||
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to("cpu",
|
||||
dtype=get_torch_dtype(
|
||||
'fp16'))
|
||||
save_file(state_dict, self.rescale_config.prompt_tensors)
|
||||
|
||||
self.print("Encoding complete.")
|
||||
|
||||
# move to cpu to save vram
|
||||
# We don't need text encoder anymore, but keep it on cpu for sampling
|
||||
# if text encoder is list
|
||||
# Move train model encoder to cpu
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder in self.sd.text_encoder:
|
||||
encoder.to("cpu")
|
||||
encoder.to('cpu')
|
||||
encoder.eval()
|
||||
encoder.requires_grad_(False)
|
||||
else:
|
||||
self.sd.text_encoder.to("cpu")
|
||||
self.prompt_cache = cache
|
||||
self.sd.text_encoder.to('cpu')
|
||||
self.sd.text_encoder.eval()
|
||||
self.sd.text_encoder.requires_grad_(False)
|
||||
|
||||
# self.sd.unet.to('cpu')
|
||||
flush()
|
||||
|
||||
self.get_latent_tensors()
|
||||
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# get random encoded prompt from cache
|
||||
prompt_txt = self.prompt_txt_list[
|
||||
torch.randint(0, len(self.prompt_txt_list), (1,)).item()
|
||||
]
|
||||
prompt = self.prompt_cache[prompt_txt].to(device=self.device_torch, dtype=dtype)
|
||||
prompt.text_embeds.to(device=self.device_torch, dtype=dtype)
|
||||
neutral = self.prompt_cache[""].to(device=self.device_torch, dtype=dtype)
|
||||
neutral.text_embeds.to(device=self.device_torch, dtype=dtype)
|
||||
if hasattr(prompt, 'pooled_embeds') \
|
||||
and hasattr(neutral, 'pooled_embeds') \
|
||||
and prompt.pooled_embeds is not None \
|
||||
and neutral.pooled_embeds is not None:
|
||||
prompt.pooled_embeds.to(device=self.device_torch, dtype=dtype)
|
||||
neutral.pooled_embeds.to(device=self.device_torch, dtype=dtype)
|
||||
|
||||
if prompt is None:
|
||||
raise ValueError(f"Prompt {prompt_txt} is not in cache")
|
||||
|
||||
loss_function = torch.nn.MSELoss()
|
||||
|
||||
with torch.no_grad():
|
||||
# self.sd.noise_scheduler.set_timesteps(
|
||||
# self.train_config.max_denoising_steps, device=self.device_torch
|
||||
# )
|
||||
# train it
|
||||
# Begin gradient accumulation
|
||||
self.sd.unet.train()
|
||||
self.sd.unet.requires_grad_(True)
|
||||
self.sd.unet.to(self.device_torch, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
# # ger a random number of steps
|
||||
timesteps_to = torch.randint(
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
).item()
|
||||
# pick random latent tensor
|
||||
latent_path = random.choice(self.latent_paths)
|
||||
latent_tensor = load_file(latent_path)
|
||||
|
||||
# get noise
|
||||
latents = self.get_latent_noise(
|
||||
pixel_height=self.rescale_config.from_resolution,
|
||||
pixel_width=self.rescale_config.from_resolution,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
noise_pred_target = (latent_tensor['noise_pred_target']).to(self.device_torch, dtype=dtype)
|
||||
latents = (latent_tensor['latents']).to(self.device_torch, dtype=dtype)
|
||||
guidance_scale = (latent_tensor['guidance_scale']).item()
|
||||
timestep = int((latent_tensor['timestep']).item())
|
||||
timesteps_to = int((latent_tensor['timesteps_to']).item())
|
||||
# seed = int((latent_tensor['seed']).item())
|
||||
|
||||
self.sd.pipeline.to(self.device_torch)
|
||||
torch.set_default_device(self.device_torch)
|
||||
|
||||
# turn off progress bar
|
||||
self.sd.pipeline.set_progress_bar_config(disable=True)
|
||||
|
||||
# get random guidance scale from 1.0 to 10.0
|
||||
guidance_scale = torch.rand(1).item() * 9.0 + 1.0
|
||||
|
||||
loss_arr = []
|
||||
|
||||
|
||||
max_len_timestep_str = len(str(self.train_config.max_denoising_steps))
|
||||
# pad with spaces
|
||||
timestep_str = str(timesteps_to).rjust(max_len_timestep_str, " ")
|
||||
new_description = f"{self.job.name} ts: {timestep_str}"
|
||||
self.progress_bar.set_description(new_description)
|
||||
|
||||
def pre_condition_callback(target_pred, input_latents):
|
||||
# handle any manipulations before feeding to our network
|
||||
reduced_pred = self.reduce_size_fn(target_pred)
|
||||
reduced_latents = self.reduce_size_fn(input_latents)
|
||||
self.optimizer.zero_grad()
|
||||
return reduced_pred, reduced_latents
|
||||
|
||||
def each_step_callback(noise_target, noise_train_pred):
|
||||
noise_target.requires_grad = False
|
||||
loss = loss_function(noise_target, noise_train_pred)
|
||||
loss_arr.append(loss.item())
|
||||
loss.backward()
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
# run the pipeline
|
||||
self.sd.pipeline.transfer_diffuse(
|
||||
num_inference_steps=timesteps_to,
|
||||
latents=latents,
|
||||
prompt_embeds=prompt.text_embeds,
|
||||
negative_prompt_embeds=neutral.text_embeds,
|
||||
pooled_prompt_embeds=prompt.pooled_embeds,
|
||||
negative_pooled_prompt_embeds=neutral.pooled_embeds,
|
||||
output_type="latent",
|
||||
num_images_per_prompt=self.train_config.batch_size,
|
||||
guidance_scale=guidance_scale,
|
||||
network=self.network,
|
||||
target_unet=self.sd.unet,
|
||||
pre_condition_callback=pre_condition_callback,
|
||||
each_step_callback=each_step_callback,
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
self.empty_embedding, # unconditional (negative prompt)
|
||||
self.empty_embedding, # conditional (positive prompt)
|
||||
self.train_config.batch_size,
|
||||
)
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
timesteps_to, device=self.device_torch
|
||||
)
|
||||
|
||||
denoised_target = self.sd.noise_scheduler.step(noise_pred_target, timestep, latents).prev_sample
|
||||
|
||||
# get the reduced latents
|
||||
# reduced_pred = self.reduce_size_fn(noise_pred_target.detach())
|
||||
denoised_target = self.reduce_size_fn(denoised_target.detach())
|
||||
reduced_latents = self.reduce_size_fn(latents.detach())
|
||||
|
||||
denoised_target.requires_grad = False
|
||||
self.optimizer.zero_grad()
|
||||
noise_pred_train = self.sd.predict_noise(
|
||||
reduced_latents,
|
||||
text_embeddings=text_embeddings,
|
||||
timestep=timestep,
|
||||
guidance_scale=guidance_scale
|
||||
)
|
||||
denoised_pred = self.sd.noise_scheduler.step(noise_pred_train, timestep, reduced_latents).prev_sample
|
||||
loss = loss_function(denoised_pred, denoised_target)
|
||||
loss_float = loss.item()
|
||||
loss.backward()
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
flush()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
# average losses
|
||||
s = 0
|
||||
for num in loss_arr:
|
||||
s += num
|
||||
|
||||
avg_loss = s / len(loss_arr)
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': avg_loss},
|
||||
{'loss': loss_float},
|
||||
)
|
||||
|
||||
return loss_dict
|
||||
|
||||
@@ -1,34 +1,31 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import random
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Optional, Union
|
||||
|
||||
from safetensors.torch import save_file, load_file
|
||||
import torch.utils.checkpoint as cp
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.layers import CheckpointGradients
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
from toolkit.prompt_utils import \
|
||||
EncodedPromptPair, ACTION_TYPES_SLIDER, \
|
||||
EncodedAnchor, concat_prompt_pairs, \
|
||||
concat_anchors, PromptEmbedsCache, encode_prompts_to_cache, build_prompt_pair_batch_from_cache, split_anchors, \
|
||||
split_prompt_pairs
|
||||
|
||||
import torch
|
||||
from leco import train_util, model_util
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
|
||||
|
||||
|
||||
class ACTION_TYPES_SLIDER:
|
||||
ERASE_NEGATIVE = 0
|
||||
ENHANCE_NEGATIVE = 1
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
@@ -36,73 +33,6 @@ def flush():
|
||||
gc.collect()
|
||||
|
||||
|
||||
class EncodedPromptPair:
|
||||
def __init__(
|
||||
self,
|
||||
target_class,
|
||||
target_class_with_neutral,
|
||||
positive_target,
|
||||
positive_target_with_neutral,
|
||||
negative_target,
|
||||
negative_target_with_neutral,
|
||||
neutral,
|
||||
empty_prompt,
|
||||
both_targets,
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
multiplier=1.0,
|
||||
weight=1.0
|
||||
):
|
||||
self.target_class = target_class
|
||||
self.target_class_with_neutral = target_class_with_neutral
|
||||
self.positive_target = positive_target
|
||||
self.positive_target_with_neutral = positive_target_with_neutral
|
||||
self.negative_target = negative_target
|
||||
self.negative_target_with_neutral = negative_target_with_neutral
|
||||
self.neutral = neutral
|
||||
self.empty_prompt = empty_prompt
|
||||
self.both_targets = both_targets
|
||||
self.multiplier = multiplier
|
||||
self.action: int = action
|
||||
self.weight = weight
|
||||
|
||||
# simulate torch to for tensors
|
||||
def to(self, *args, **kwargs):
|
||||
self.target_class = self.target_class.to(*args, **kwargs)
|
||||
self.positive_target = self.positive_target.to(*args, **kwargs)
|
||||
self.positive_target_with_neutral = self.positive_target_with_neutral.to(*args, **kwargs)
|
||||
self.negative_target = self.negative_target.to(*args, **kwargs)
|
||||
self.negative_target_with_neutral = self.negative_target_with_neutral.to(*args, **kwargs)
|
||||
self.neutral = self.neutral.to(*args, **kwargs)
|
||||
self.empty_prompt = self.empty_prompt.to(*args, **kwargs)
|
||||
self.both_targets = self.both_targets.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
|
||||
class PromptEmbedsCache:
|
||||
prompts: dict[str, PromptEmbeds] = {}
|
||||
|
||||
def __setitem__(self, __name: str, __value: PromptEmbeds) -> None:
|
||||
self.prompts[__name] = __value
|
||||
|
||||
def __getitem__(self, __name: str) -> Optional[PromptEmbeds]:
|
||||
if __name in self.prompts:
|
||||
return self.prompts[__name]
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
class EncodedAnchor:
|
||||
def __init__(
|
||||
self,
|
||||
prompt,
|
||||
neg_prompt,
|
||||
multiplier=1.0
|
||||
):
|
||||
self.prompt = prompt
|
||||
self.neg_prompt = neg_prompt
|
||||
self.multiplier = multiplier
|
||||
|
||||
|
||||
class TrainSliderProcess(BaseSDTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
@@ -115,24 +45,26 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
self.prompt_cache = PromptEmbedsCache()
|
||||
self.prompt_pairs: list[EncodedPromptPair] = []
|
||||
self.anchor_pairs: list[EncodedAnchor] = []
|
||||
# keep track of prompt chunk size
|
||||
self.prompt_chunk_size = 1
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.print(f"Loading prompt file from {self.slider_config.prompt_file}")
|
||||
|
||||
# read line by line from file
|
||||
if self.slider_config.prompt_file:
|
||||
with open(self.slider_config.prompt_file, 'r') as f:
|
||||
self.print(f"Loading prompt file from {self.slider_config.prompt_file}")
|
||||
with open(self.slider_config.prompt_file, 'r', encoding='utf-8') as f:
|
||||
self.prompt_txt_list = f.readlines()
|
||||
# clean empty lines
|
||||
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
|
||||
|
||||
self.print(f"Loaded {len(self.prompt_txt_list)} prompts. Encoding them..")
|
||||
|
||||
self.print(f"Found {len(self.prompt_txt_list)} prompts.")
|
||||
|
||||
if not self.slider_config.prompt_tensors:
|
||||
print(f"Prompt tensors not found. Building prompt tensors for {self.train_config.steps} steps.")
|
||||
# shuffle
|
||||
random.shuffle(self.prompt_txt_list)
|
||||
# trim to max steps
|
||||
@@ -143,163 +75,57 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
|
||||
# get encoded latents for our prompts
|
||||
with torch.no_grad():
|
||||
if self.slider_config.prompt_tensors is not None:
|
||||
# check to see if it exists
|
||||
if os.path.exists(self.slider_config.prompt_tensors):
|
||||
# load it.
|
||||
self.print(f"Loading prompt tensors from {self.slider_config.prompt_tensors}")
|
||||
prompt_tensors = load_file(self.slider_config.prompt_tensors, device='cpu')
|
||||
# add them to the cache
|
||||
for prompt_txt, prompt_tensor in tqdm(prompt_tensors.items(), desc="Loading prompts", leave=False):
|
||||
if prompt_txt.startswith("te:"):
|
||||
prompt = prompt_txt[3:]
|
||||
# text_embeds
|
||||
text_embeds = prompt_tensor
|
||||
pooled_embeds = None
|
||||
# find pool embeds
|
||||
if f"pe:{prompt}" in prompt_tensors:
|
||||
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
|
||||
# list of neutrals. Can come from file or be empty
|
||||
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
|
||||
|
||||
# make it
|
||||
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
|
||||
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
|
||||
# build the prompts to cache
|
||||
prompts_to_cache = []
|
||||
for neutral in neutral_list:
|
||||
for target in self.slider_config.targets:
|
||||
prompt_list = [
|
||||
f"{target.target_class}", # target_class
|
||||
f"{target.target_class} {neutral}", # target_class with neutral
|
||||
f"{target.positive}", # positive_target
|
||||
f"{target.positive} {neutral}", # positive_target with neutral
|
||||
f"{target.negative}", # negative_target
|
||||
f"{target.negative} {neutral}", # negative_target with neutral
|
||||
f"{neutral}", # neutral
|
||||
f"{target.positive} {target.negative}", # both targets
|
||||
f"{target.negative} {target.positive}", # both targets reverse
|
||||
]
|
||||
prompts_to_cache += prompt_list
|
||||
|
||||
if len(cache.prompts) == 0:
|
||||
print("Prompt tensors not found. Encoding prompts..")
|
||||
empty_prompt = ""
|
||||
# encode empty_prompt
|
||||
cache[empty_prompt] = self.sd.encode_prompt(empty_prompt)
|
||||
# remove duplicates
|
||||
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
|
||||
|
||||
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
|
||||
|
||||
for neutral in tqdm(neutral_list, desc="Encoding prompts", leave=False):
|
||||
for target in self.slider_config.targets:
|
||||
prompt_list = [
|
||||
f"{target.target_class}", # target_class
|
||||
f"{target.target_class} {neutral}", # target_class with neutral
|
||||
f"{target.positive}", # positive_target
|
||||
f"{target.positive} {neutral}", # positive_target with neutral
|
||||
f"{target.negative}", # negative_target
|
||||
f"{target.negative} {neutral}", # negative_target with neutral
|
||||
f"{neutral}", # neutral
|
||||
f"{target.positive} {target.negative}", # both targets
|
||||
f"{target.negative} {target.positive}", # both targets
|
||||
]
|
||||
for p in prompt_list:
|
||||
# build the cache
|
||||
if cache[p] is None:
|
||||
cache[p] = self.sd.encode_prompt(p).to(device="cpu", dtype=torch.float32)
|
||||
|
||||
erase_negative = len(target.positive.strip()) == 0
|
||||
enhance_positive = len(target.negative.strip()) == 0
|
||||
|
||||
both = not erase_negative and not enhance_positive
|
||||
|
||||
if erase_negative and enhance_positive:
|
||||
raise ValueError("target must have at least one of positive or negative or both")
|
||||
# for slider we need to have an enhancer, an eraser, and then
|
||||
# an inverse with negative weights to balance the network
|
||||
# if we don't do this, we will get different contrast and focus.
|
||||
# we only perform actions of enhancing and erasing on the negative
|
||||
# todo work on way to do all of this in one shot
|
||||
if self.slider_config.prompt_tensors:
|
||||
print(f"Saving prompt tensors to {self.slider_config.prompt_tensors}")
|
||||
state_dict = {}
|
||||
for prompt_txt, prompt_embeds in cache.prompts.items():
|
||||
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to("cpu",
|
||||
dtype=get_torch_dtype('fp16'))
|
||||
if prompt_embeds.pooled_embeds is not None:
|
||||
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to("cpu",
|
||||
dtype=get_torch_dtype(
|
||||
'fp16'))
|
||||
save_file(state_dict, self.slider_config.prompt_tensors)
|
||||
# encode them
|
||||
cache = encode_prompts_to_cache(
|
||||
prompt_list=prompts_to_cache,
|
||||
sd=self.sd,
|
||||
cache=cache,
|
||||
prompt_tensor_file=self.slider_config.prompt_tensors
|
||||
)
|
||||
|
||||
prompt_pairs = []
|
||||
for neutral in tqdm(neutral_list, desc="Encoding prompts", leave=False):
|
||||
prompt_batches = []
|
||||
for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
|
||||
for target in self.slider_config.targets:
|
||||
erase_negative = len(target.positive.strip()) == 0
|
||||
enhance_positive = len(target.negative.strip()) == 0
|
||||
prompt_pair_batch = build_prompt_pair_batch_from_cache(
|
||||
cache=cache,
|
||||
target=target,
|
||||
neutral=neutral,
|
||||
|
||||
both = not erase_negative and not enhance_positive
|
||||
|
||||
if both or erase_negative:
|
||||
print("Encoding erase negative")
|
||||
prompt_pairs += [
|
||||
# erase standard
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.positive}"],
|
||||
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
negative_target=cache[f"{target.negative}"],
|
||||
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
multiplier=target.multiplier,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
if both or enhance_positive:
|
||||
print("Encoding enhance positive")
|
||||
prompt_pairs += [
|
||||
# enhance standard, swap pos neg
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.negative}"],
|
||||
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
negative_target=cache[f"{target.positive}"],
|
||||
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||
multiplier=target.multiplier,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
# if both or enhance_positive:
|
||||
if both:
|
||||
print("Encoding erase positive (inverse)")
|
||||
prompt_pairs += [
|
||||
# erase inverted
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.negative}"],
|
||||
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
negative_target=cache[f"{target.positive}"],
|
||||
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
multiplier=target.multiplier * -1.0,
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
# if both or erase_negative:
|
||||
if both:
|
||||
print("Encoding enhance negative (inverse)")
|
||||
prompt_pairs += [
|
||||
# enhance inverted
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.positive}"],
|
||||
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
negative_target=cache[f"{target.negative}"],
|
||||
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||
empty_prompt=cache[""],
|
||||
multiplier=target.multiplier * -1.0,
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
)
|
||||
if self.slider_config.batch_full_slide:
|
||||
# concat the prompt pairs
|
||||
# this allows us to run the entire 4 part process in one shot (for slider)
|
||||
self.prompt_chunk_size = 4
|
||||
concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
|
||||
prompt_pairs += [concat_prompt_pair_batch]
|
||||
else:
|
||||
self.prompt_chunk_size = 1
|
||||
# do them one at a time (probably not necessary after new optimizations)
|
||||
prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
|
||||
|
||||
# setup anchors
|
||||
anchor_pairs = []
|
||||
@@ -312,14 +138,26 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
if cache[prompt] == None:
|
||||
cache[prompt] = self.sd.encode_prompt(prompt)
|
||||
|
||||
anchor_pairs += [
|
||||
EncodedAnchor(
|
||||
prompt=cache[anchor.prompt],
|
||||
neg_prompt=cache[anchor.neg_prompt],
|
||||
multiplier=anchor.multiplier
|
||||
)
|
||||
]
|
||||
anchor_batch = []
|
||||
# we get the prompt pair multiplier from first prompt pair
|
||||
# since they are all the same. We need to match their network polarity
|
||||
prompt_pair_multipliers = prompt_pairs[0].multiplier_list
|
||||
for prompt_multiplier in prompt_pair_multipliers:
|
||||
# match the network multiplier polarity
|
||||
anchor_scalar = 1.0 if prompt_multiplier > 0 else -1.0
|
||||
anchor_batch += [
|
||||
EncodedAnchor(
|
||||
prompt=cache[anchor.prompt],
|
||||
neg_prompt=cache[anchor.neg_prompt],
|
||||
multiplier=anchor.multiplier * anchor_scalar
|
||||
)
|
||||
]
|
||||
|
||||
anchor_pairs += [
|
||||
concat_anchors(anchor_batch).to('cpu')
|
||||
]
|
||||
if len(anchor_pairs) > 0:
|
||||
self.anchor_pairs = anchor_pairs
|
||||
|
||||
# move to cpu to save vram
|
||||
# We don't need text encoder anymore, but keep it on cpu for sampling
|
||||
@@ -331,17 +169,13 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
self.sd.text_encoder.to("cpu")
|
||||
self.prompt_cache = cache
|
||||
self.prompt_pairs = prompt_pairs
|
||||
self.anchor_pairs = anchor_pairs
|
||||
# self.anchor_pairs = anchor_pairs
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# get random multiplier between 1 and 3
|
||||
rand_weight = 1
|
||||
# rand_weight = torch.rand((1,)).item() * 2 + 1
|
||||
|
||||
# get a random pair
|
||||
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
||||
torch.randint(0, len(self.prompt_pairs), (1,)).item()
|
||||
@@ -353,18 +187,17 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
height, width = self.slider_config.resolutions[
|
||||
torch.randint(0, len(self.slider_config.resolutions), (1,)).item()
|
||||
]
|
||||
if self.train_config.gradient_checkpointing:
|
||||
# may get disabled elsewhere
|
||||
self.sd.unet.enable_gradient_checkpointing()
|
||||
|
||||
weight = prompt_pair.weight
|
||||
multiplier = prompt_pair.multiplier
|
||||
|
||||
unet = self.sd.unet
|
||||
noise_scheduler = self.sd.noise_scheduler
|
||||
optimizer = self.optimizer
|
||||
lr_scheduler = self.lr_scheduler
|
||||
loss_function = torch.nn.MSELoss()
|
||||
|
||||
def get_noise_pred(neg, pos, gs, cts, dn):
|
||||
return self.predict_noise(
|
||||
return self.sd.predict_noise(
|
||||
latents=dn,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
neg, # negative prompt
|
||||
@@ -375,9 +208,6 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
guidance_scale=gs,
|
||||
)
|
||||
|
||||
# set network multiplier
|
||||
self.network.multiplier = multiplier * rand_weight
|
||||
|
||||
with torch.no_grad():
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
@@ -390,10 +220,15 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
).item()
|
||||
|
||||
# for a complete slider, the batch size is 4 to begin with now
|
||||
true_batch_size = prompt_pair.target_class.text_embeds.shape[0] * self.train_config.batch_size
|
||||
|
||||
# get noise
|
||||
noise = self.get_latent_noise(
|
||||
noise = self.sd.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
batch_size=true_batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get latents
|
||||
@@ -402,8 +237,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
self.network.multiplier = multiplier * rand_weight
|
||||
denoised_latents = self.diffuse_some_steps(
|
||||
# pass the multiplier list to the network
|
||||
self.network.multiplier = prompt_pair.multiplier_list
|
||||
denoised_latents = self.sd.diffuse_some_steps(
|
||||
latents, # pass simple noise latents
|
||||
train_tools.concat_prompt_embeddings(
|
||||
prompt_pair.positive_target, # unconditional
|
||||
@@ -415,19 +251,27 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
guidance_scale=3,
|
||||
)
|
||||
|
||||
# split the latents into out prompt pair chunks
|
||||
denoised_latent_chunks = torch.chunk(denoised_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
noise_scheduler.set_timesteps(1000)
|
||||
|
||||
current_timestep = noise_scheduler.timesteps[
|
||||
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
||||
]
|
||||
|
||||
# flush() # 4.2GB to 3GB on 512x512
|
||||
|
||||
# 4.20 GB RAM for 512x512
|
||||
positive_latents = get_noise_pred(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.negative_target, # positive prompt
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
)
|
||||
positive_latents.requires_grad = False
|
||||
positive_latents_chunks = torch.chunk(positive_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
neutral_latents = get_noise_pred(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
@@ -435,7 +279,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
)
|
||||
neutral_latents.requires_grad = False
|
||||
neutral_latents_chunks = torch.chunk(neutral_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
unconditional_latents = get_noise_pred(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
@@ -443,87 +289,142 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
|
||||
anchor_loss = None
|
||||
if len(self.anchor_pairs) > 0:
|
||||
# get a random anchor pair
|
||||
anchor: EncodedAnchor = self.anchor_pairs[
|
||||
torch.randint(0, len(self.anchor_pairs), (1,)).item()
|
||||
]
|
||||
with torch.no_grad():
|
||||
anchor_target_noise = get_noise_pred(
|
||||
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
with self.network:
|
||||
# anchor whatever weight prompt pair is using
|
||||
pos_nem_mult = 1.0 if prompt_pair.multiplier > 0 else -1.0
|
||||
self.network.multiplier = anchor.multiplier * pos_nem_mult * rand_weight
|
||||
|
||||
anchor_pred_noise = get_noise_pred(
|
||||
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
|
||||
self.network.multiplier = prompt_pair.multiplier * rand_weight
|
||||
|
||||
with self.network:
|
||||
self.network.multiplier = prompt_pair.multiplier * rand_weight
|
||||
target_latents = get_noise_pred(
|
||||
prompt_pair.positive_target,
|
||||
prompt_pair.target_class,
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
|
||||
# if self.logging_config.verbose:
|
||||
# self.print("target_latents:", target_latents[0, 0, :5, :5])
|
||||
|
||||
positive_latents.requires_grad = False
|
||||
neutral_latents.requires_grad = False
|
||||
unconditional_latents.requires_grad = False
|
||||
if len(self.anchor_pairs) > 0:
|
||||
anchor_target_noise.requires_grad = False
|
||||
anchor_loss = loss_function(
|
||||
anchor_target_noise,
|
||||
anchor_pred_noise,
|
||||
)
|
||||
erase = prompt_pair.action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE
|
||||
guidance_scale = 1.0
|
||||
unconditional_latents.requires_grad = False
|
||||
unconditional_latents_chunks = torch.chunk(unconditional_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
offset = guidance_scale * (positive_latents - unconditional_latents)
|
||||
flush() # 4.2GB to 3GB on 512x512
|
||||
|
||||
offset_neutral = neutral_latents
|
||||
if erase:
|
||||
offset_neutral -= offset
|
||||
else:
|
||||
# enhance
|
||||
offset_neutral += offset
|
||||
# 4.20 GB RAM for 512x512
|
||||
anchor_loss_float = None
|
||||
if len(self.anchor_pairs) > 0:
|
||||
with torch.no_grad():
|
||||
# get a random anchor pair
|
||||
anchor: EncodedAnchor = self.anchor_pairs[
|
||||
torch.randint(0, len(self.anchor_pairs), (1,)).item()
|
||||
]
|
||||
anchor.to(self.device_torch, dtype=dtype)
|
||||
|
||||
loss = loss_function(
|
||||
target_latents,
|
||||
offset_neutral,
|
||||
) * weight
|
||||
# first we get the target prediction without network active
|
||||
anchor_target_noise = get_noise_pred(
|
||||
anchor.neg_prompt, anchor.prompt, 1, current_timestep, denoised_latents
|
||||
# ).to("cpu", dtype=torch.float32)
|
||||
).requires_grad_(False)
|
||||
|
||||
loss_slide = loss.item()
|
||||
# to save vram, we will run these through separately while tracking grads
|
||||
# otherwise it consumes a ton of vram and this isn't our speed bottleneck
|
||||
anchor_chunks = split_anchors(anchor, self.prompt_chunk_size)
|
||||
anchor_target_noise_chunks = torch.chunk(anchor_target_noise, self.prompt_chunk_size, dim=0)
|
||||
assert len(anchor_chunks) == len(denoised_latent_chunks)
|
||||
|
||||
if anchor_loss is not None:
|
||||
loss += anchor_loss
|
||||
# 4.32 GB RAM for 512x512
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
anchor_float_losses = []
|
||||
for anchor_chunk, denoised_latent_chunk, anchor_target_noise_chunk in zip(
|
||||
anchor_chunks, denoised_latent_chunks, anchor_target_noise_chunks
|
||||
):
|
||||
self.network.multiplier = anchor_chunk.multiplier_list
|
||||
|
||||
loss_float = loss.item()
|
||||
anchor_pred_noise = get_noise_pred(
|
||||
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
|
||||
)
|
||||
# 9.42 GB RAM for 512x512 -> 4.20 GB RAM for 512x512 with new grad_checkpointing
|
||||
anchor_loss = loss_function(
|
||||
anchor_target_noise_chunk,
|
||||
anchor_pred_noise,
|
||||
)
|
||||
anchor_float_losses.append(anchor_loss.item())
|
||||
# compute anchor loss gradients
|
||||
# we will accumulate them later
|
||||
# this saves a ton of memory doing them separately
|
||||
anchor_loss.backward()
|
||||
del anchor_pred_noise
|
||||
del anchor_target_noise_chunk
|
||||
del anchor_loss
|
||||
flush()
|
||||
|
||||
loss = loss.to(self.device_torch)
|
||||
anchor_loss_float = sum(anchor_float_losses) / len(anchor_float_losses)
|
||||
del anchor_chunks
|
||||
del anchor_target_noise_chunks
|
||||
del anchor_target_noise
|
||||
# move anchor back to cpu
|
||||
anchor.to("cpu")
|
||||
flush()
|
||||
|
||||
prompt_pair_chunks = split_prompt_pairs(prompt_pair, self.prompt_chunk_size)
|
||||
assert len(prompt_pair_chunks) == len(denoised_latent_chunks)
|
||||
# 3.28 GB RAM for 512x512
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
loss_list = []
|
||||
for prompt_pair_chunk, \
|
||||
denoised_latent_chunk, \
|
||||
positive_latents_chunk, \
|
||||
neutral_latents_chunk, \
|
||||
unconditional_latents_chunk \
|
||||
in zip(
|
||||
prompt_pair_chunks,
|
||||
denoised_latent_chunks,
|
||||
positive_latents_chunks,
|
||||
neutral_latents_chunks,
|
||||
unconditional_latents_chunks,
|
||||
):
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list
|
||||
target_latents = get_noise_pred(
|
||||
prompt_pair_chunk.positive_target,
|
||||
prompt_pair_chunk.target_class,
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latent_chunk
|
||||
)
|
||||
|
||||
guidance_scale = 1.0
|
||||
|
||||
offset = guidance_scale * (positive_latents_chunk - unconditional_latents_chunk)
|
||||
|
||||
# make offset multiplier based on actions
|
||||
offset_multiplier_list = []
|
||||
for action in prompt_pair_chunk.action_list:
|
||||
if action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE:
|
||||
offset_multiplier_list += [-1.0]
|
||||
elif action == ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE:
|
||||
offset_multiplier_list += [1.0]
|
||||
|
||||
offset_multiplier = torch.tensor(offset_multiplier_list).to(offset.device, dtype=offset.dtype)
|
||||
# make offset multiplier match rank of offset
|
||||
offset_multiplier = offset_multiplier.view(offset.shape[0], 1, 1, 1)
|
||||
offset *= offset_multiplier
|
||||
|
||||
offset_neutral = neutral_latents_chunk
|
||||
# offsets are already adjusted on a per-batch basis
|
||||
offset_neutral += offset
|
||||
|
||||
# 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
|
||||
loss = loss_function(
|
||||
target_latents,
|
||||
offset_neutral,
|
||||
) * prompt_pair_chunk.weight
|
||||
|
||||
loss.backward()
|
||||
loss_list.append(loss.item())
|
||||
del target_latents
|
||||
del offset_neutral
|
||||
del loss
|
||||
flush()
|
||||
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
loss_float = sum(loss_list) / len(loss_list)
|
||||
if anchor_loss_float is not None:
|
||||
loss_float += anchor_loss_float
|
||||
|
||||
del (
|
||||
positive_latents,
|
||||
neutral_latents,
|
||||
unconditional_latents,
|
||||
target_latents,
|
||||
latents,
|
||||
latents
|
||||
)
|
||||
# move back to cpu
|
||||
prompt_pair.to("cpu")
|
||||
@@ -535,9 +436,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss_float},
|
||||
)
|
||||
if anchor_loss is not None:
|
||||
loss_dict['sl_l'] = loss_slide
|
||||
loss_dict['an_l'] = anchor_loss.item()
|
||||
if anchor_loss_float is not None:
|
||||
loss_dict['sl_l'] = loss_float
|
||||
loss_dict['an_l'] = anchor_loss_float
|
||||
|
||||
return loss_dict
|
||||
# end hook_train_loop
|
||||
|
||||
@@ -221,7 +221,7 @@ class TrainSliderProcessOld(BaseSDTrainProcess):
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# get a random pair
|
||||
@@ -245,7 +245,7 @@ class TrainSliderProcessOld(BaseSDTrainProcess):
|
||||
loss_function = torch.nn.MSELoss()
|
||||
|
||||
def get_noise_pred(p, n, gs, cts, dn):
|
||||
return self.predict_noise(
|
||||
return self.sd.predict_noise(
|
||||
latents=dn,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
p, # unconditional
|
||||
@@ -272,9 +272,11 @@ class TrainSliderProcessOld(BaseSDTrainProcess):
|
||||
).item()
|
||||
|
||||
# get noise
|
||||
noise = self.get_latent_noise(
|
||||
noise = self.sd.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
batch_size=self.train_config.batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get latents
|
||||
@@ -284,7 +286,7 @@ class TrainSliderProcessOld(BaseSDTrainProcess):
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
self.network.multiplier = multiplier
|
||||
denoised_latents = self.diffuse_some_steps(
|
||||
denoised_latents = self.sd.diffuse_some_steps(
|
||||
latents, # pass simple noise latents
|
||||
train_tools.concat_prompt_embeddings(
|
||||
positive, # unconditional
|
||||
|
||||
@@ -24,6 +24,7 @@ from diffusers import AutoencoderKL
|
||||
from tqdm import tqdm
|
||||
import time
|
||||
import numpy as np
|
||||
from .models.vgg19_critic import Critic
|
||||
|
||||
IMAGE_TRANSFORMS = transforms.Compose(
|
||||
[
|
||||
@@ -37,145 +38,6 @@ def unnormalize(tensor):
|
||||
return (tensor / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
|
||||
class Critic:
|
||||
process: 'TrainVAEProcess'
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
learning_rate=1e-5,
|
||||
device='cpu',
|
||||
optimizer='adam',
|
||||
num_critic_per_gen=1,
|
||||
dtype='float32',
|
||||
lambda_gp=10,
|
||||
start_step=0,
|
||||
warmup_steps=1000,
|
||||
process=None,
|
||||
optimizer_params=None,
|
||||
):
|
||||
self.learning_rate = learning_rate
|
||||
self.device = device
|
||||
self.optimizer_type = optimizer
|
||||
self.num_critic_per_gen = num_critic_per_gen
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.process = process
|
||||
self.model = None
|
||||
self.optimizer = None
|
||||
self.scheduler = None
|
||||
self.warmup_steps = warmup_steps
|
||||
self.start_step = start_step
|
||||
self.lambda_gp = lambda_gp
|
||||
|
||||
if optimizer_params is None:
|
||||
optimizer_params = {}
|
||||
self.optimizer_params = optimizer_params
|
||||
self.print = self.process.print
|
||||
print(f" Critic config: {self.__dict__}")
|
||||
|
||||
def setup(self):
|
||||
from .models.vgg19_critic import Vgg19Critic
|
||||
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
|
||||
self.load_weights()
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
params = self.model.parameters()
|
||||
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
|
||||
self.optimizer,
|
||||
total_iters=self.process.max_steps * self.num_critic_per_gen,
|
||||
factor=1,
|
||||
verbose=False
|
||||
)
|
||||
|
||||
def load_weights(self):
|
||||
path_to_load = None
|
||||
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
|
||||
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
path_to_load = latest_file
|
||||
else:
|
||||
self.print(f" - No checkpoint found, starting from scratch")
|
||||
if path_to_load:
|
||||
self.model.load_state_dict(load_file(path_to_load))
|
||||
|
||||
def save(self, step=None):
|
||||
self.process.update_training_metadata()
|
||||
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zeropad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
|
||||
save_file(self.model.state_dict(), save_path, save_meta)
|
||||
self.print(f"Saved critic to {save_path}")
|
||||
|
||||
def get_critic_loss(self, vgg_output):
|
||||
if self.start_step > self.process.step_num:
|
||||
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
|
||||
|
||||
warmup_scaler = 1.0
|
||||
# we need a warmup when we come on of 1000 steps
|
||||
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
|
||||
if self.process.step_num < self.start_step + self.warmup_steps:
|
||||
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
|
||||
# set model to not train for generator loss
|
||||
self.model.eval()
|
||||
self.model.requires_grad_(False)
|
||||
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
|
||||
|
||||
# run model
|
||||
stacked_output = self.model(vgg_pred)
|
||||
|
||||
return (-torch.mean(stacked_output)) * warmup_scaler
|
||||
|
||||
def step(self, vgg_output):
|
||||
|
||||
# train critic here
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
|
||||
critic_losses = []
|
||||
for i in range(self.num_critic_per_gen):
|
||||
inputs = vgg_output.detach()
|
||||
inputs = inputs.to(self.device, dtype=self.torch_dtype)
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
|
||||
|
||||
stacked_output = self.model(inputs)
|
||||
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
|
||||
|
||||
# Compute gradient penalty
|
||||
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
|
||||
|
||||
# Compute WGAN-GP critic loss
|
||||
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
|
||||
critic_loss.backward()
|
||||
self.optimizer.zero_grad()
|
||||
self.optimizer.step()
|
||||
self.scheduler.step()
|
||||
critic_losses.append(critic_loss.item())
|
||||
|
||||
# avg loss
|
||||
loss = np.mean(critic_losses)
|
||||
return loss
|
||||
|
||||
def get_lr(self):
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
learning_rate = (
|
||||
self.optimizer.param_groups[0]["d"] *
|
||||
self.optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
learning_rate = self.optimizer.param_groups[0]['lr']
|
||||
|
||||
return learning_rate
|
||||
|
||||
|
||||
class TrainVAEProcess(BaseTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
@@ -10,3 +10,7 @@ from .TrainSliderProcessOld import TrainSliderProcessOld
|
||||
from .TrainLoRAHack import TrainLoRAHack
|
||||
from .TrainSDRescaleProcess import TrainSDRescaleProcess
|
||||
from .ModRescaleLoraProcess import ModRescaleLoraProcess
|
||||
from .GenerateProcess import GenerateProcess
|
||||
from .BaseExtensionProcess import BaseExtensionProcess
|
||||
from .TrainESRGANProcess import TrainESRGANProcess
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess
|
||||
|
||||
@@ -1,5 +1,17 @@
|
||||
import glob
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
from toolkit.losses import get_gradient_penalty
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.optimizer import get_optimizer
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
|
||||
class MeanReduce(nn.Module):
|
||||
@@ -36,3 +48,147 @@ class Vgg19Critic(nn.Module):
|
||||
|
||||
def forward(self, inputs):
|
||||
return self.main(inputs)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jobs.process.TrainVAEProcess import TrainVAEProcess
|
||||
from jobs.process.TrainESRGANProcess import TrainESRGANProcess
|
||||
|
||||
|
||||
class Critic:
|
||||
process: Union['TrainVAEProcess', 'TrainESRGANProcess']
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
learning_rate=1e-5,
|
||||
device='cpu',
|
||||
optimizer='adam',
|
||||
num_critic_per_gen=1,
|
||||
dtype='float32',
|
||||
lambda_gp=10,
|
||||
start_step=0,
|
||||
warmup_steps=1000,
|
||||
process=None,
|
||||
optimizer_params=None,
|
||||
):
|
||||
self.learning_rate = learning_rate
|
||||
self.device = device
|
||||
self.optimizer_type = optimizer
|
||||
self.num_critic_per_gen = num_critic_per_gen
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.process = process
|
||||
self.model = None
|
||||
self.optimizer = None
|
||||
self.scheduler = None
|
||||
self.warmup_steps = warmup_steps
|
||||
self.start_step = start_step
|
||||
self.lambda_gp = lambda_gp
|
||||
|
||||
if optimizer_params is None:
|
||||
optimizer_params = {}
|
||||
self.optimizer_params = optimizer_params
|
||||
self.print = self.process.print
|
||||
print(f" Critic config: {self.__dict__}")
|
||||
|
||||
def setup(self):
|
||||
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
|
||||
self.load_weights()
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
params = self.model.parameters()
|
||||
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
|
||||
self.optimizer,
|
||||
total_iters=self.process.max_steps * self.num_critic_per_gen,
|
||||
factor=1,
|
||||
verbose=False
|
||||
)
|
||||
|
||||
def load_weights(self):
|
||||
path_to_load = None
|
||||
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
|
||||
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
path_to_load = latest_file
|
||||
else:
|
||||
self.print(f" - No checkpoint found, starting from scratch")
|
||||
if path_to_load:
|
||||
self.model.load_state_dict(load_file(path_to_load))
|
||||
|
||||
def save(self, step=None):
|
||||
self.process.update_training_metadata()
|
||||
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zeropad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
|
||||
save_file(self.model.state_dict(), save_path, save_meta)
|
||||
self.print(f"Saved critic to {save_path}")
|
||||
|
||||
def get_critic_loss(self, vgg_output):
|
||||
if self.start_step > self.process.step_num:
|
||||
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
|
||||
|
||||
warmup_scaler = 1.0
|
||||
# we need a warmup when we come on of 1000 steps
|
||||
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
|
||||
if self.process.step_num < self.start_step + self.warmup_steps:
|
||||
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
|
||||
# set model to not train for generator loss
|
||||
self.model.eval()
|
||||
self.model.requires_grad_(False)
|
||||
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
|
||||
|
||||
# run model
|
||||
stacked_output = self.model(vgg_pred)
|
||||
|
||||
return (-torch.mean(stacked_output)) * warmup_scaler
|
||||
|
||||
def step(self, vgg_output):
|
||||
|
||||
# train critic here
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
|
||||
critic_losses = []
|
||||
for i in range(self.num_critic_per_gen):
|
||||
inputs = vgg_output.detach()
|
||||
inputs = inputs.to(self.device, dtype=self.torch_dtype)
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
|
||||
|
||||
stacked_output = self.model(inputs)
|
||||
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
|
||||
|
||||
# Compute gradient penalty
|
||||
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
|
||||
|
||||
# Compute WGAN-GP critic loss
|
||||
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
|
||||
critic_loss.backward()
|
||||
self.optimizer.zero_grad()
|
||||
self.optimizer.step()
|
||||
self.scheduler.step()
|
||||
critic_losses.append(critic_loss.item())
|
||||
|
||||
# avg loss
|
||||
loss = np.mean(critic_losses)
|
||||
return loss
|
||||
|
||||
def get_lr(self):
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
learning_rate = (
|
||||
self.optimizer.param_groups[0]["d"] *
|
||||
self.optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
learning_rate = self.optimizer.param_groups[0]['lr']
|
||||
|
||||
return learning_rate
|
||||
|
||||
|
||||
338
notebooks/SliderTraining.ipynb
Normal file
338
notebooks/SliderTraining.ipynb
Normal file
@@ -0,0 +1,338 @@
|
||||
{
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"provenance": [],
|
||||
"machine_shape": "hm",
|
||||
"gpuType": "V100"
|
||||
},
|
||||
"kernelspec": {
|
||||
"name": "python3",
|
||||
"display_name": "Python 3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
},
|
||||
"accelerator": "GPU"
|
||||
},
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## Slider Training\n",
|
||||
"\n",
|
||||
"This is a quick colab demo for training sliders like can be found in my CivitAI profile https://civitai.com/user/Ostris/models . I will work on making it more user friendly, but for now, it will get you started."
|
||||
],
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image"
|
||||
],
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
],
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict({\n",
|
||||
" # This is the config I use on my sliders, It is solid and tested\n",
|
||||
" 'job': 'train',\n",
|
||||
" 'config': {\n",
|
||||
" # the name will be used to create a folder in the output folder\n",
|
||||
" # it will also replace any [name] token in the rest of this config\n",
|
||||
" 'name': 'detail_slider_v1',\n",
|
||||
" # folder will be created with name above in folder below\n",
|
||||
" # it can be relative to the project root or absolute\n",
|
||||
" 'training_folder': \"output/LoRA\",\n",
|
||||
" 'device': 'cuda', # cpu, cuda:0, etc\n",
|
||||
" # for tensorboard logging, we will make a subfolder for this job\n",
|
||||
" 'log_dir': \"output/.tensorboard\",\n",
|
||||
" # you can stack processes for other jobs, It is not tested with sliders though\n",
|
||||
" # just use one for now\n",
|
||||
" 'process': [\n",
|
||||
" {\n",
|
||||
" 'type': 'slider', # tells runner to run the slider process\n",
|
||||
" # network is the LoRA network for a slider, I recommend to leave this be\n",
|
||||
" 'network': {\n",
|
||||
" 'type': \"lora\",\n",
|
||||
" # rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good\n",
|
||||
" 'linear': 8, # \"rank\" or \"dim\"\n",
|
||||
" 'linear_alpha': 4, # Do about half of rank \"alpha\"\n",
|
||||
" # 'conv': 4, # for convolutional layers \"locon\"\n",
|
||||
" # 'conv_alpha': 4, # Do about half of conv \"alpha\"\n",
|
||||
" },\n",
|
||||
" # training config\n",
|
||||
" 'train': {\n",
|
||||
" # this is also used in sampling. Stick with ddpm unless you know what you are doing\n",
|
||||
" 'noise_scheduler': \"ddpm\", # or \"ddpm\", \"lms\", \"euler_a\"\n",
|
||||
" # how many steps to train. More is not always better. I rarely go over 1000\n",
|
||||
" 'steps': 100,\n",
|
||||
" # I have had good results with 4e-4 to 1e-4 at 500 steps\n",
|
||||
" 'lr': 2e-4,\n",
|
||||
" # enables gradient checkpoint, saves vram, leave it on\n",
|
||||
" 'gradient_checkpointing': True,\n",
|
||||
" # train the unet. I recommend leaving this true\n",
|
||||
" 'train_unet': True,\n",
|
||||
" # train the text encoder. I don't recommend this unless you have a special use case\n",
|
||||
" # for sliders we are adjusting representation of the concept (unet),\n",
|
||||
" # not the description of it (text encoder)\n",
|
||||
" 'train_text_encoder': False,\n",
|
||||
"\n",
|
||||
" # just leave unless you know what you are doing\n",
|
||||
" # also supports \"dadaptation\" but set lr to 1 if you use that,\n",
|
||||
" # but it learns too fast and I don't recommend it\n",
|
||||
" 'optimizer': \"adamw\",\n",
|
||||
" # only constant for now\n",
|
||||
" 'lr_scheduler': \"constant\",\n",
|
||||
" # we randomly denoise random num of steps form 1 to this number\n",
|
||||
" # while training. Just leave it\n",
|
||||
" 'max_denoising_steps': 40,\n",
|
||||
" # works great at 1. I do 1 even with my 4090.\n",
|
||||
" # higher may not work right with newer single batch stacking code anyway\n",
|
||||
" 'batch_size': 1,\n",
|
||||
" # bf16 works best if your GPU supports it (modern)\n",
|
||||
" 'dtype': 'bf16', # fp32, bf16, fp16\n",
|
||||
" # I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX\n",
|
||||
" # although, the way we train sliders is comparative, so it probably won't work anyway\n",
|
||||
" 'noise_offset': 0.0,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # the model to train the LoRA network on\n",
|
||||
" 'model': {\n",
|
||||
" # name_or_path can be a hugging face name, local path or url to model\n",
|
||||
" # on civit ai with or without modelVersionId. They will be cached in /model folder\n",
|
||||
" # epicRealisim v5\n",
|
||||
" 'name_or_path': \"https://civitai.com/models/25694?modelVersionId=134065\",\n",
|
||||
" 'is_v2': False, # for v2 models\n",
|
||||
" 'is_v_pred': False, # for v-prediction models (most v2 models)\n",
|
||||
" # has some issues with the dual text encoder and the way we train sliders\n",
|
||||
" # it works bit weights need to probably be higher to see it.\n",
|
||||
" 'is_xl': False, # for SDXL models\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # saving config\n",
|
||||
" 'save': {\n",
|
||||
" 'dtype': 'float16', # precision to save. I recommend float16\n",
|
||||
" 'save_every': 50, # save every this many steps\n",
|
||||
" # this will remove step counts more than this number\n",
|
||||
" # allows you to save more often in case of a crash without filling up your drive\n",
|
||||
" 'max_step_saves_to_keep': 2,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # sampling config\n",
|
||||
" 'sample': {\n",
|
||||
" # must match train.noise_scheduler, this is not used here\n",
|
||||
" # but may be in future and in other processes\n",
|
||||
" 'sampler': \"ddpm\",\n",
|
||||
" # sample every this many steps\n",
|
||||
" 'sample_every': 20,\n",
|
||||
" # image size\n",
|
||||
" 'width': 512,\n",
|
||||
" 'height': 512,\n",
|
||||
" # prompts to use for sampling. Do as many as you want, but it slows down training\n",
|
||||
" # pick ones that will best represent the concept you are trying to adjust\n",
|
||||
" # allows some flags after the prompt\n",
|
||||
" # --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive\n",
|
||||
" # slide are good tests. will inherit sample.network_multiplier if not set\n",
|
||||
" # --n [string] # negative prompt, will inherit sample.neg if not set\n",
|
||||
" # Only 75 tokens allowed currently\n",
|
||||
" # I like to do a wide positive and negative spread so I can see a good range and stop\n",
|
||||
" # early if the network is braking down\n",
|
||||
" 'prompts': [\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5\",\n",
|
||||
" \"a golden retriever sitting on a leather couch, --m -5\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m -3\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m 3\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m 5\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5\",\n",
|
||||
" ],\n",
|
||||
" # negative prompt used on all prompts above as default if they don't have one\n",
|
||||
" 'neg': \"cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome\",\n",
|
||||
" # seed for sampling. 42 is the answer for everything\n",
|
||||
" 'seed': 42,\n",
|
||||
" # walks the seed so s1 is 42, s2 is 43, s3 is 44, etc\n",
|
||||
" # will start over on next sample_every so s1 is always seed\n",
|
||||
" # works well if you use same prompt but want different results\n",
|
||||
" 'walk_seed': False,\n",
|
||||
" # cfg scale (4 to 10 is good)\n",
|
||||
" 'guidance_scale': 7,\n",
|
||||
" # sampler steps (20 to 30 is good)\n",
|
||||
" 'sample_steps': 20,\n",
|
||||
" # default network multiplier for all prompts\n",
|
||||
" # since we are training a slider, I recommend overriding this with --m [number]\n",
|
||||
" # in the prompts above to get both sides of the slider\n",
|
||||
" 'network_multiplier': 1.0,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # logging information\n",
|
||||
" 'logging': {\n",
|
||||
" 'log_every': 10, # log every this many steps\n",
|
||||
" 'use_wandb': False, # not supported yet\n",
|
||||
" 'verbose': False, # probably done need unless you are debugging\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # slider training config, best for last\n",
|
||||
" 'slider': {\n",
|
||||
" # resolutions to train on. [ width, height ]. This is less important for sliders\n",
|
||||
" # as we are not teaching the model anything it doesn't already know\n",
|
||||
" # but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1\n",
|
||||
" # and [ 1024, 1024 ] for sd_xl\n",
|
||||
" # you can do as many as you want here\n",
|
||||
" 'resolutions': [\n",
|
||||
" [512, 512],\n",
|
||||
" # [ 512, 768 ]\n",
|
||||
" # [ 768, 768 ]\n",
|
||||
" ],\n",
|
||||
" # slider training uses 4 combined steps for a single round. This will do it in one gradient\n",
|
||||
" # step. It is highly optimized and shouldn't take anymore vram than doing without it,\n",
|
||||
" # since we break down batches for gradient accumulation now. so just leave it on.\n",
|
||||
" 'batch_full_slide': True,\n",
|
||||
" # These are the concepts to train on. You can do as many as you want here,\n",
|
||||
" # but they can conflict outweigh each other. Other than experimenting, I recommend\n",
|
||||
" # just doing one for good results\n",
|
||||
" 'targets': [\n",
|
||||
" # target_class is the base concept we are adjusting the representation of\n",
|
||||
" # for example, if we are adjusting the representation of a person, we would use \"person\"\n",
|
||||
" # if we are adjusting the representation of a cat, we would use \"cat\" It is not\n",
|
||||
" # a keyword necessarily but what the model understands the concept to represent.\n",
|
||||
" # \"person\" will affect men, women, children, etc but will not affect cats, dogs, etc\n",
|
||||
" # it is the models base general understanding of the concept and everything it represents\n",
|
||||
" # you can leave it blank to affect everything. In this example, we are adjusting\n",
|
||||
" # detail, so we will leave it blank to affect everything\n",
|
||||
" {\n",
|
||||
" 'target_class': \"\",\n",
|
||||
" # positive is the prompt for the positive side of the slider.\n",
|
||||
" # It is the concept that will be excited and amplified in the model when we slide the slider\n",
|
||||
" # to the positive side and forgotten / inverted when we slide\n",
|
||||
" # the slider to the negative side. It is generally best to include the target_class in\n",
|
||||
" # the prompt. You want it to be the extreme of what you want to train on. For example,\n",
|
||||
" # if you want to train on fat people, you would use \"an extremely fat, morbidly obese person\"\n",
|
||||
" # as the prompt. Not just \"fat person\"\n",
|
||||
" # max 75 tokens for now\n",
|
||||
" 'positive': \"high detail, 8k, intricate, detailed, high resolution, high res, high quality\",\n",
|
||||
" # negative is the prompt for the negative side of the slider and works the same as positive\n",
|
||||
" # it does not necessarily work the same as a negative prompt when generating images\n",
|
||||
" # these need to be polar opposites.\n",
|
||||
" # max 76 tokens for now\n",
|
||||
" 'negative': \"blurry, boring, fuzzy, low detail, low resolution, low res, low quality\",\n",
|
||||
" # the loss for this target is multiplied by this number.\n",
|
||||
" # if you are doing more than one target it may be good to set less important ones\n",
|
||||
" # to a lower number like 0.1 so they don't outweigh the primary target\n",
|
||||
" 'weight': 1.0,\n",
|
||||
" },\n",
|
||||
" ],\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
" ]\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # You can put any information you want here, and it will be saved in the model.\n",
|
||||
" # The below is an example, but you can put your grocery list in it if you want.\n",
|
||||
" # It is saved in the model so be aware of that. The software will include this\n",
|
||||
" # plus some other information for you automatically\n",
|
||||
" 'meta': {\n",
|
||||
" # [name] gets replaced with the name above\n",
|
||||
" 'name': \"[name]\",\n",
|
||||
" 'version': '1.0',\n",
|
||||
" # 'creator': {\n",
|
||||
" # 'name': 'your name',\n",
|
||||
" # 'email': 'your@gmail.com',\n",
|
||||
" # 'website': 'https://your.website'\n",
|
||||
" # }\n",
|
||||
" }\n",
|
||||
"})\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. Ill update soon."
|
||||
],
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -5,7 +5,6 @@ diffusers
|
||||
transformers
|
||||
lycoris_lora
|
||||
flatten_json
|
||||
accelerator
|
||||
pyyaml
|
||||
oyaml
|
||||
tensorboard
|
||||
@@ -14,4 +13,6 @@ invisible-watermark
|
||||
einops
|
||||
accelerate
|
||||
toml
|
||||
albumentations
|
||||
albumentations
|
||||
pydantic
|
||||
omegaconf
|
||||
|
||||
2
run.py
2
run.py
@@ -1,5 +1,7 @@
|
||||
import os
|
||||
import sys
|
||||
from typing import Union, OrderedDict
|
||||
|
||||
sys.path.insert(0, os.getcwd())
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
|
||||
@@ -2,6 +2,7 @@ import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
from diffusers.loaders import LoraLoaderMixin
|
||||
from safetensors.torch import load_file
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
@@ -63,8 +64,8 @@ keys_in_both.sort()
|
||||
|
||||
json_data = {
|
||||
"both": keys_in_both,
|
||||
"state_dict_2": keys_not_in_state_dict_2,
|
||||
"state_dict_1": keys_not_in_state_dict_1
|
||||
"not_in_state_dict_2": keys_not_in_state_dict_2,
|
||||
"not_in_state_dict_1": keys_not_in_state_dict_1
|
||||
}
|
||||
json_data = json.dumps(json_data, indent=4)
|
||||
|
||||
@@ -84,6 +85,15 @@ project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
json_save_path = os.path.join(project_root, 'config', 'keys.json')
|
||||
json_matched_save_path = os.path.join(project_root, 'config', 'matched.json')
|
||||
json_duped_save_path = os.path.join(project_root, 'config', 'duped.json')
|
||||
state_dict_1_filename = os.path.basename(args.file_1[0])
|
||||
state_dict_2_filename = os.path.basename(args.file_2[0])
|
||||
# save key names for each in own file
|
||||
with open(os.path.join(project_root, 'config', f'{state_dict_1_filename}.json'), 'w') as f:
|
||||
f.write(json.dumps(state_dict_1_keys, indent=4))
|
||||
|
||||
with open(os.path.join(project_root, 'config', f'{state_dict_2_filename}.json'), 'w') as f:
|
||||
f.write(json.dumps(state_dict_2_keys, indent=4))
|
||||
|
||||
|
||||
with open(json_save_path, 'w') as f:
|
||||
f.write(json_data)
|
||||
4
toolkit/basic.py
Normal file
4
toolkit/basic.py
Normal file
@@ -0,0 +1,4 @@
|
||||
|
||||
|
||||
def value_map(inputs, min_in, max_in, min_out, max_out):
|
||||
return (inputs - min_in) * (max_out - min_out) / (max_in - min_in) + min_out
|
||||
217
toolkit/civitai.py
Normal file
217
toolkit/civitai.py
Normal file
@@ -0,0 +1,217 @@
|
||||
from toolkit.paths import MODELS_PATH
|
||||
import requests
|
||||
import os
|
||||
import json
|
||||
import tqdm
|
||||
|
||||
|
||||
class ModelCache:
|
||||
def __init__(self):
|
||||
self.raw_cache = {}
|
||||
self.cache_path = os.path.join(MODELS_PATH, '.ai_toolkit_cache.json')
|
||||
if os.path.exists(self.cache_path):
|
||||
with open(self.cache_path, 'r') as f:
|
||||
all_cache = json.load(f)
|
||||
if 'models' in all_cache:
|
||||
self.raw_cache = all_cache['models']
|
||||
else:
|
||||
self.raw_cache = all_cache
|
||||
|
||||
def get_model_path(self, model_id: int, model_version_id: int = None):
|
||||
if str(model_id) not in self.raw_cache:
|
||||
return None
|
||||
if model_version_id is None:
|
||||
# get latest version
|
||||
model_version_id = max([int(x) for x in self.raw_cache[str(model_id)].keys()])
|
||||
if model_version_id is None:
|
||||
return None
|
||||
model_path = self.raw_cache[str(model_id)][str(model_version_id)]['model_path']
|
||||
# check if model path exists
|
||||
if not os.path.exists(model_path):
|
||||
# remove version from cache
|
||||
del self.raw_cache[str(model_id)][str(model_version_id)]
|
||||
self.save()
|
||||
return None
|
||||
return model_path
|
||||
else:
|
||||
if str(model_version_id) not in self.raw_cache[str(model_id)]:
|
||||
return None
|
||||
model_path = self.raw_cache[str(model_id)][str(model_version_id)]['model_path']
|
||||
# check if model path exists
|
||||
if not os.path.exists(model_path):
|
||||
# remove version from cache
|
||||
del self.raw_cache[str(model_id)][str(model_version_id)]
|
||||
self.save()
|
||||
return None
|
||||
return model_path
|
||||
|
||||
def update_cache(self, model_id: int, model_version_id: int, model_path: str):
|
||||
if str(model_id) not in self.raw_cache:
|
||||
self.raw_cache[str(model_id)] = {}
|
||||
if str(model_version_id) not in self.raw_cache[str(model_id)]:
|
||||
self.raw_cache[str(model_id)][str(model_version_id)] = {}
|
||||
self.raw_cache[str(model_id)][str(model_version_id)] = {
|
||||
'model_path': model_path
|
||||
}
|
||||
self.save()
|
||||
|
||||
def save(self):
|
||||
if not os.path.exists(os.path.dirname(self.cache_path)):
|
||||
os.makedirs(os.path.dirname(self.cache_path), exist_ok=True)
|
||||
all_cache = {'models': {}}
|
||||
if os.path.exists(self.cache_path):
|
||||
# load it first
|
||||
with open(self.cache_path, 'r') as f:
|
||||
all_cache = json.load(f)
|
||||
|
||||
all_cache['models'] = self.raw_cache
|
||||
|
||||
with open(self.cache_path, 'w') as f:
|
||||
json.dump(all_cache, f, indent=2)
|
||||
|
||||
|
||||
def get_model_download_info(model_id: int, model_version_id: int = None):
|
||||
# curl https://civitai.com/api/v1/models?limit=3&types=TextualInversion \
|
||||
# -H "Content-Type: application/json" \
|
||||
# -X GET
|
||||
print(
|
||||
f"Getting model info for model id: {model_id}{f' and version id: {model_version_id}' if model_version_id is not None else ''}")
|
||||
endpoint = f"https://civitai.com/api/v1/models/{model_id}"
|
||||
|
||||
# get the json
|
||||
response = requests.get(endpoint)
|
||||
response.raise_for_status()
|
||||
model_data = response.json()
|
||||
|
||||
model_version = None
|
||||
|
||||
# go through versions and get the top one if one is not set
|
||||
for version in model_data['modelVersions']:
|
||||
if model_version_id is not None:
|
||||
if str(version['id']) == str(model_version_id):
|
||||
model_version = version
|
||||
break
|
||||
else:
|
||||
# get first version
|
||||
model_version = version
|
||||
break
|
||||
|
||||
if model_version is None:
|
||||
raise ValueError(
|
||||
f"Could not find a model version for model id: {model_id}{f' and version id: {model_version_id}' if model_version_id is not None else ''}")
|
||||
|
||||
model_file = None
|
||||
# go through files and prefer fp16 safetensors
|
||||
# "metadata": {
|
||||
# "fp": "fp16",
|
||||
# "size": "pruned",
|
||||
# "format": "SafeTensor"
|
||||
# },
|
||||
# todo check pickle scans and skip if not good
|
||||
# try to get fp16 safetensor
|
||||
for file in model_version['files']:
|
||||
if file['metadata']['fp'] == 'fp16' and file['metadata']['format'] == 'SafeTensor':
|
||||
model_file = file
|
||||
break
|
||||
|
||||
if model_file is None:
|
||||
# try to get primary
|
||||
for file in model_version['files']:
|
||||
if file['primary']:
|
||||
model_file = file
|
||||
break
|
||||
|
||||
if model_file is None:
|
||||
# try to get any safetensor
|
||||
for file in model_version['files']:
|
||||
if file['metadata']['format'] == 'SafeTensor':
|
||||
model_file = file
|
||||
break
|
||||
|
||||
if model_file is None:
|
||||
# try to get any fp16
|
||||
for file in model_version['files']:
|
||||
if file['metadata']['fp'] == 'fp16':
|
||||
model_file = file
|
||||
break
|
||||
|
||||
if model_file is None:
|
||||
# try to get any
|
||||
for file in model_version['files']:
|
||||
model_file = file
|
||||
break
|
||||
|
||||
if model_file is None:
|
||||
raise ValueError(f"Could not find a model file to download for model id: {model_id}")
|
||||
|
||||
return model_file, model_version['id']
|
||||
|
||||
|
||||
def get_model_path_from_url(url: str):
|
||||
# get query params form url if they are set
|
||||
# https: // civitai.com / models / 25694?modelVersionId = 127742
|
||||
query_params = {}
|
||||
if '?' in url:
|
||||
query_string = url.split('?')[1]
|
||||
query_params = dict(qc.split("=") for qc in query_string.split("&"))
|
||||
|
||||
# get model id from url
|
||||
model_id = url.split('/')[-1]
|
||||
# remove query params from model id
|
||||
if '?' in model_id:
|
||||
model_id = model_id.split('?')[0]
|
||||
if model_id.isdigit():
|
||||
model_id = int(model_id)
|
||||
else:
|
||||
raise ValueError(f"Invalid model id: {model_id}")
|
||||
|
||||
model_cache = ModelCache()
|
||||
model_path = model_cache.get_model_path(model_id, query_params.get('modelVersionId', None))
|
||||
if model_path is not None:
|
||||
return model_path
|
||||
else:
|
||||
# download model
|
||||
file_info, model_version_id = get_model_download_info(model_id, query_params.get('modelVersionId', None))
|
||||
|
||||
download_url = file_info['downloadUrl'] # url does not work directly
|
||||
size_kb = file_info['sizeKB']
|
||||
filename = file_info['name']
|
||||
model_path = os.path.join(MODELS_PATH, filename)
|
||||
|
||||
# download model
|
||||
print(f"Did not find model locally, downloading from model from: {download_url}")
|
||||
|
||||
# use tqdm to show status of downlod
|
||||
response = requests.get(download_url, stream=True)
|
||||
response.raise_for_status()
|
||||
total_size_in_bytes = int(response.headers.get('content-length', 0))
|
||||
block_size = 1024 # 1 Kibibyte
|
||||
progress_bar = tqdm.tqdm(total=total_size_in_bytes, unit='iB', unit_scale=True)
|
||||
tmp_path = os.path.join(MODELS_PATH, f".download_tmp_{filename}")
|
||||
os.makedirs(os.path.dirname(model_path), exist_ok=True)
|
||||
# remove tmp file if it exists
|
||||
if os.path.exists(tmp_path):
|
||||
os.remove(tmp_path)
|
||||
|
||||
try:
|
||||
|
||||
with open(tmp_path, 'wb') as f:
|
||||
for data in response.iter_content(block_size):
|
||||
progress_bar.update(len(data))
|
||||
f.write(data)
|
||||
progress_bar.close()
|
||||
# move to final path
|
||||
os.rename(tmp_path, model_path)
|
||||
model_cache.update_cache(model_id, model_version_id, model_path)
|
||||
|
||||
return model_path
|
||||
except Exception as e:
|
||||
# remove tmp file
|
||||
os.remove(tmp_path)
|
||||
raise e
|
||||
|
||||
|
||||
# if is main
|
||||
if __name__ == '__main__':
|
||||
model_path = get_model_path_from_url("https://civitai.com/models/25694?modelVersionId=127742")
|
||||
print(model_path)
|
||||
@@ -1,5 +1,7 @@
|
||||
import os
|
||||
import json
|
||||
from typing import Union
|
||||
|
||||
import oyaml as yaml
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
@@ -47,7 +49,17 @@ fixed_loader.add_implicit_resolver(
|
||||
list(u'-+0123456789.'))
|
||||
|
||||
|
||||
def get_config(config_file_path, name=None):
|
||||
def get_config(
|
||||
config_file_path_or_dict: Union[str, dict, OrderedDict],
|
||||
name=None
|
||||
):
|
||||
# if we got a dict, process it and return it
|
||||
if isinstance(config_file_path_or_dict, dict) or isinstance(config_file_path_or_dict, OrderedDict):
|
||||
config = config_file_path_or_dict
|
||||
return preprocess_config(config, name)
|
||||
|
||||
config_file_path = config_file_path_or_dict
|
||||
|
||||
# first check if it is in the config folder
|
||||
config_path = os.path.join(TOOLKIT_ROOT, 'config', config_file_path)
|
||||
# see if it is in the config folder with any of the possible extensions if it doesnt have one
|
||||
@@ -70,10 +82,10 @@ def get_config(config_file_path, name=None):
|
||||
|
||||
# if we found it, check if it is a json or yaml file
|
||||
if real_config_path.endswith('.json') or real_config_path.endswith('.jsonc'):
|
||||
with open(real_config_path, 'r') as f:
|
||||
with open(real_config_path, 'r', encoding='utf-8') as f:
|
||||
config = json.load(f, object_pairs_hook=OrderedDict)
|
||||
elif real_config_path.endswith('.yaml') or real_config_path.endswith('.yml'):
|
||||
with open(real_config_path, 'r') as f:
|
||||
with open(real_config_path, 'r', encoding='utf-8') as f:
|
||||
config = yaml.load(f, Loader=fixed_loader)
|
||||
else:
|
||||
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
from typing import List
|
||||
import os
|
||||
import time
|
||||
from typing import List, Optional
|
||||
import random
|
||||
|
||||
|
||||
class SaveConfig:
|
||||
@@ -27,6 +30,7 @@ class SampleConfig:
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.network_multiplier = kwargs.get('network_multiplier', 1)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
|
||||
|
||||
class NetworkConfig:
|
||||
@@ -35,13 +39,15 @@ class NetworkConfig:
|
||||
rank = kwargs.get('rank', None)
|
||||
linear = kwargs.get('linear', None)
|
||||
if rank is not None:
|
||||
self.rank: int = rank # rank for backward compatibility
|
||||
self.rank: int = rank # rank for backward compatibility
|
||||
self.linear: int = rank
|
||||
elif linear is not None:
|
||||
self.rank: int = linear
|
||||
self.linear: int = linear
|
||||
self.conv: int = kwargs.get('conv', None)
|
||||
self.alpha: float = kwargs.get('alpha', 1.0)
|
||||
self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha)
|
||||
self.conv_alpha: float = kwargs.get('conv_alpha', self.conv)
|
||||
|
||||
|
||||
class TrainConfig:
|
||||
@@ -60,7 +66,7 @@ class TrainConfig:
|
||||
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
||||
self.optimizer_params = kwargs.get('optimizer_params', {})
|
||||
self.skip_first_sample = kwargs.get('skip_first_sample', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True)
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
@@ -69,6 +75,8 @@ class ModelConfig:
|
||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.vae_path: str = kwargs.get('vae_path', None)
|
||||
|
||||
if self.name_or_path is None:
|
||||
raise ValueError('name_or_path must be specified')
|
||||
@@ -101,3 +109,198 @@ class SliderConfig:
|
||||
self.resolutions: List[List[int]] = kwargs.get('resolutions', [[512, 512]])
|
||||
self.prompt_file: str = kwargs.get('prompt_file', None)
|
||||
self.prompt_tensors: str = kwargs.get('prompt_tensors', None)
|
||||
self.batch_full_slide: bool = kwargs.get('batch_full_slide', True)
|
||||
|
||||
|
||||
class GenerateImageConfig:
|
||||
def __init__(
|
||||
self,
|
||||
prompt: str = '',
|
||||
prompt_2: Optional[str] = None,
|
||||
width: int = 512,
|
||||
height: int = 512,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: str = '',
|
||||
negative_prompt_2: Optional[str] = None,
|
||||
seed: int = -1,
|
||||
network_multiplier: float = 1.0,
|
||||
guidance_rescale: float = 0.0,
|
||||
# the tag [time] will be replaced with milliseconds since epoch
|
||||
output_path: str = None, # full image path
|
||||
output_folder: str = None, # folder to save image in if output_path is not specified
|
||||
output_ext: str = 'png', # extension to save image as if output_path is not specified
|
||||
output_tail: str = '', # tail to add to output filename
|
||||
add_prompt_file: bool = False, # add a prompt file with generated image
|
||||
):
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
self.num_inference_steps: int = num_inference_steps
|
||||
self.guidance_scale: float = guidance_scale
|
||||
self.guidance_rescale: float = guidance_rescale
|
||||
self.prompt: str = prompt
|
||||
self.prompt_2: str = prompt_2
|
||||
self.negative_prompt: str = negative_prompt
|
||||
self.negative_prompt_2: str = negative_prompt_2
|
||||
|
||||
self.output_path: str = output_path
|
||||
self.seed: int = seed
|
||||
if self.seed == -1:
|
||||
# generate random one
|
||||
self.seed = random.randint(0, 2 ** 32 - 1)
|
||||
self.network_multiplier: float = network_multiplier
|
||||
self.output_folder: str = output_folder
|
||||
self.output_ext: str = output_ext
|
||||
self.add_prompt_file: bool = add_prompt_file
|
||||
self.output_tail: str = output_tail
|
||||
self.gen_time: int = int(time.time() * 1000)
|
||||
|
||||
# prompt string will override any settings above
|
||||
self._process_prompt_string()
|
||||
|
||||
# handle dual text encoder prompts if nothing passed
|
||||
if negative_prompt_2 is None:
|
||||
self.negative_prompt_2 = negative_prompt
|
||||
|
||||
if prompt_2 is None:
|
||||
self.prompt_2 = prompt
|
||||
|
||||
# parse prompt paths
|
||||
if self.output_path is None and self.output_folder is None:
|
||||
raise ValueError('output_path or output_folder must be specified')
|
||||
elif self.output_path is not None:
|
||||
self.output_folder = os.path.dirname(self.output_path)
|
||||
self.output_ext = os.path.splitext(self.output_path)[1][1:]
|
||||
self.output_filename_no_ext = os.path.splitext(os.path.basename(self.output_path))[0]
|
||||
|
||||
else:
|
||||
self.output_filename_no_ext = '[time]_[count]'
|
||||
if len(self.output_tail) > 0:
|
||||
self.output_filename_no_ext += '_' + self.output_tail
|
||||
self.output_path = os.path.join(self.output_folder, self.output_filename_no_ext + '.' + self.output_ext)
|
||||
|
||||
# adjust height
|
||||
self.height = max(64, self.height - self.height % 8) # round to divisible by 8
|
||||
self.width = max(64, self.width - self.width % 8) # round to divisible by 8
|
||||
|
||||
def set_gen_time(self, gen_time: int = None):
|
||||
if gen_time is not None:
|
||||
self.gen_time = gen_time
|
||||
else:
|
||||
self.gen_time = int(time.time() * 1000)
|
||||
|
||||
def _get_path_no_ext(self, count: int = 0, max_count=0):
|
||||
# zero pad count
|
||||
count_str = str(count).zfill(len(str(max_count)))
|
||||
# replace [time] with gen time
|
||||
filename = self.output_filename_no_ext.replace('[time]', str(self.gen_time))
|
||||
# replace [count] with count
|
||||
filename = filename.replace('[count]', count_str)
|
||||
return filename
|
||||
|
||||
def get_image_path(self, count: int = 0, max_count=0):
|
||||
filename = self._get_path_no_ext(count, max_count)
|
||||
filename += '.' + self.output_ext
|
||||
# join with folder
|
||||
return os.path.join(self.output_folder, filename)
|
||||
|
||||
def get_prompt_path(self, count: int = 0, max_count=0):
|
||||
filename = self._get_path_no_ext(count, max_count)
|
||||
filename += '.txt'
|
||||
# join with folder
|
||||
return os.path.join(self.output_folder, filename)
|
||||
|
||||
def save_image(self, image, count: int = 0, max_count=0):
|
||||
# make parent dirs
|
||||
os.makedirs(self.output_folder, exist_ok=True)
|
||||
self.set_gen_time()
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
|
||||
def save_prompt_file(self, count: int = 0, max_count=0):
|
||||
# save prompt file
|
||||
with open(self.get_prompt_path(count, max_count), 'w') as f:
|
||||
prompt = self.prompt
|
||||
if self.prompt_2 is not None:
|
||||
prompt += ' --p2 ' + self.prompt_2
|
||||
if self.negative_prompt is not None:
|
||||
prompt += ' --n ' + self.negative_prompt
|
||||
if self.negative_prompt_2 is not None:
|
||||
prompt += ' --n2 ' + self.negative_prompt_2
|
||||
prompt += ' --w ' + str(self.width)
|
||||
prompt += ' --h ' + str(self.height)
|
||||
prompt += ' --seed ' + str(self.seed)
|
||||
prompt += ' --cfg ' + str(self.guidance_scale)
|
||||
prompt += ' --steps ' + str(self.num_inference_steps)
|
||||
prompt += ' --m ' + str(self.network_multiplier)
|
||||
prompt += ' --gr ' + str(self.guidance_rescale)
|
||||
|
||||
# get gen info
|
||||
f.write(self.prompt)
|
||||
|
||||
def _process_prompt_string(self):
|
||||
# we will try to support all sd-scripts where we can
|
||||
|
||||
# FROM SD-SCRIPTS
|
||||
# --n Treat everything until the next option as a negative prompt.
|
||||
# --w Specify the width of the generated image.
|
||||
# --h Specify the height of the generated image.
|
||||
# --d Specify the seed for the generated image.
|
||||
# --l Specify the CFG scale for the generated image.
|
||||
# --s Specify the number of steps during generation.
|
||||
|
||||
# OURS and some QOL additions
|
||||
# --m Specify the network multiplier for the generated image.
|
||||
# --p2 Prompt for the second text encoder (SDXL only)
|
||||
# --n2 Negative prompt for the second text encoder (SDXL only)
|
||||
# --gr Specify the guidance rescale for the generated image (SDXL only)
|
||||
|
||||
# --seed Specify the seed for the generated image same as --d
|
||||
# --cfg Specify the CFG scale for the generated image same as --l
|
||||
# --steps Specify the number of steps during generation same as --s
|
||||
# --network_multiplier Specify the network multiplier for the generated image same as --m
|
||||
|
||||
# process prompt string and update values if it has some
|
||||
if self.prompt is not None and len(self.prompt) > 0:
|
||||
# process prompt string
|
||||
prompt = self.prompt
|
||||
prompt = prompt.strip()
|
||||
p_split = prompt.split('--')
|
||||
self.prompt = p_split[0].strip()
|
||||
|
||||
if len(p_split) > 1:
|
||||
for split in p_split[1:]:
|
||||
# allows multi char flags
|
||||
flag = split.split(' ')[0].strip()
|
||||
content = split[len(flag):].strip()
|
||||
if flag == 'p2':
|
||||
self.prompt_2 = content
|
||||
elif flag == 'n':
|
||||
self.negative_prompt = content
|
||||
elif flag == 'n2':
|
||||
self.negative_prompt_2 = content
|
||||
elif flag == 'w':
|
||||
self.width = int(content)
|
||||
elif flag == 'h':
|
||||
self.height = int(content)
|
||||
elif flag == 'd':
|
||||
self.seed = int(content)
|
||||
elif flag == 'seed':
|
||||
self.seed = int(content)
|
||||
elif flag == 'l':
|
||||
self.guidance_scale = float(content)
|
||||
elif flag == 'cfg':
|
||||
self.guidance_scale = float(content)
|
||||
elif flag == 's':
|
||||
self.num_inference_steps = int(content)
|
||||
elif flag == 'steps':
|
||||
self.num_inference_steps = int(content)
|
||||
elif flag == 'm':
|
||||
self.network_multiplier = float(content)
|
||||
elif flag == 'network_multiplier':
|
||||
self.network_multiplier = float(content)
|
||||
elif flag == 'gr':
|
||||
self.guidance_rescale = float(content)
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from torchvision import transforms
|
||||
from torch.utils.data import Dataset
|
||||
from tqdm import tqdm
|
||||
import albumentations as A
|
||||
|
||||
|
||||
class ImageDataset(Dataset):
|
||||
@@ -38,7 +42,7 @@ class ImageDataset(Dataset):
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
@@ -65,7 +69,7 @@ class ImageDataset(Dataset):
|
||||
if self.random_scale and min_img_size > self.resolution:
|
||||
if min_img_size < self.resolution:
|
||||
print(
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={file}")
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={img_path}")
|
||||
scale_size = self.resolution
|
||||
else:
|
||||
scale_size = random.randint(self.resolution, int(min_img_size))
|
||||
@@ -78,3 +82,124 @@ class ImageDataset(Dataset):
|
||||
img = self.transform(img)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
class Augments:
|
||||
def __init__(self, **kwargs):
|
||||
self.method_name = kwargs.get('method', None)
|
||||
self.params = kwargs.get('params', {})
|
||||
|
||||
# convert kwargs enums for cv2
|
||||
for key, value in self.params.items():
|
||||
if isinstance(value, str):
|
||||
# split the string
|
||||
split_string = value.split('.')
|
||||
if len(split_string) == 2 and split_string[0] == 'cv2':
|
||||
if hasattr(cv2, split_string[1]):
|
||||
self.params[key] = getattr(cv2, split_string[1].upper())
|
||||
else:
|
||||
raise ValueError(f"invalid cv2 enum: {split_string[1]}")
|
||||
|
||||
|
||||
class AugmentedImageDataset(ImageDataset):
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.augmentations = self.get_config('augmentations', [])
|
||||
self.augmentations = [Augments(**aug) for aug in self.augmentations]
|
||||
|
||||
augmentation_list = []
|
||||
for aug in self.augmentations:
|
||||
# make sure method name is valid
|
||||
assert hasattr(A, aug.method_name), f"invalid augmentation method: {aug.method_name}"
|
||||
# get the method
|
||||
method = getattr(A, aug.method_name)
|
||||
# add the method to the list
|
||||
augmentation_list.append(method(**aug.params))
|
||||
|
||||
self.aug_transform = A.Compose(augmentation_list)
|
||||
self.original_transform = self.transform
|
||||
# replace transform so we get raw pil image
|
||||
self.transform = transforms.Compose([])
|
||||
|
||||
def __getitem__(self, index):
|
||||
# get the original image
|
||||
# image is a PIL image, convert to bgr
|
||||
pil_image = super().__getitem__(index)
|
||||
open_cv_image = np.array(pil_image)
|
||||
# Convert RGB to BGR
|
||||
open_cv_image = open_cv_image[:, :, ::-1].copy()
|
||||
|
||||
# apply augmentations
|
||||
augmented = self.aug_transform(image=open_cv_image)["image"]
|
||||
|
||||
# convert back to RGB tensor
|
||||
augmented = cv2.cvtColor(augmented, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# convert to PIL image
|
||||
augmented = Image.fromarray(augmented)
|
||||
|
||||
# return both # return image as 0 - 1 tensor
|
||||
return transforms.ToTensor()(pil_image), transforms.ToTensor()(augmented)
|
||||
|
||||
|
||||
class PairedImageDataset(Dataset):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.size = self.get_config('size', 512)
|
||||
self.path = self.get_config('path', required=True)
|
||||
self.default_prompt = self.get_config('default_prompt', '')
|
||||
self.network_weight = self.get_config('network_weight', 1.0)
|
||||
self.file_list = [os.path.join(self.path, file) for file in os.listdir(self.path) if
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
def __len__(self):
|
||||
return len(self.file_list)
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
if key in self.config:
|
||||
value = self.config[key]
|
||||
return value
|
||||
elif required:
|
||||
raise ValueError(f'config file error. Missing "config.dataset.{key}" key')
|
||||
else:
|
||||
return default
|
||||
|
||||
def __getitem__(self, index):
|
||||
img_path = self.file_list[index]
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
|
||||
# see if prompt file exists
|
||||
path_no_ext = os.path.splitext(img_path)[0]
|
||||
prompt_path = path_no_ext + '.txt'
|
||||
if os.path.exists(prompt_path):
|
||||
with open(prompt_path, 'r', encoding='utf-8') as f:
|
||||
prompt = f.read()
|
||||
# remove any newlines
|
||||
prompt = prompt.replace('\n', ', ')
|
||||
# remove new lines for all operating systems
|
||||
prompt = prompt.replace('\r', ', ')
|
||||
prompt_split = prompt.split(',')
|
||||
# remove empty strings
|
||||
prompt_split = [p.strip() for p in prompt_split if p.strip()]
|
||||
# join back together
|
||||
prompt = ', '.join(prompt_split)
|
||||
else:
|
||||
prompt = self.default_prompt
|
||||
|
||||
height = self.size
|
||||
# determine width to keep aspect ratio
|
||||
width = int(img.size[0] * height / img.size[1])
|
||||
|
||||
# Downscale the source image first
|
||||
img = img.resize((width, height), Image.BICUBIC)
|
||||
img = self.transform(img)
|
||||
|
||||
return img, prompt, self.network_weight
|
||||
|
||||
|
||||
51
toolkit/esrgan_utils.py
Normal file
51
toolkit/esrgan_utils.py
Normal file
@@ -0,0 +1,51 @@
|
||||
|
||||
to_basicsr_dict = {
|
||||
'model.0.weight': 'conv_first.weight',
|
||||
'model.0.bias': 'conv_first.bias',
|
||||
'model.1.sub.23.weight': 'conv_body.weight',
|
||||
'model.1.sub.23.bias': 'conv_body.bias',
|
||||
'model.3.weight': 'conv_up1.weight',
|
||||
'model.3.bias': 'conv_up1.bias',
|
||||
'model.6.weight': 'conv_up2.weight',
|
||||
'model.6.bias': 'conv_up2.bias',
|
||||
'model.8.weight': 'conv_hr.weight',
|
||||
'model.8.bias': 'conv_hr.bias',
|
||||
'model.10.bias': 'conv_last.bias',
|
||||
'model.10.weight': 'conv_last.weight',
|
||||
# 'model.1.sub.0.RDB1.conv1.0.weight': 'body.0.rdb1.conv1.weight'
|
||||
}
|
||||
|
||||
def convert_state_dict_to_basicsr(state_dict):
|
||||
new_state_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if k in to_basicsr_dict:
|
||||
new_state_dict[to_basicsr_dict[k]] = v
|
||||
elif k.startswith('model.1.sub.'):
|
||||
bsr_name = k.replace('model.1.sub.', 'body.').lower()
|
||||
bsr_name = bsr_name.replace('.0.weight', '.weight')
|
||||
bsr_name = bsr_name.replace('.0.bias', '.bias')
|
||||
new_state_dict[bsr_name] = v
|
||||
else:
|
||||
new_state_dict[k] = v
|
||||
return new_state_dict
|
||||
|
||||
|
||||
# just matching a commonly used format
|
||||
def convert_basicsr_state_dict_to_save_format(state_dict):
|
||||
new_state_dict = {}
|
||||
to_basicsr_dict_values = list(to_basicsr_dict.values())
|
||||
for k, v in state_dict.items():
|
||||
if k in to_basicsr_dict_values:
|
||||
for key, value in to_basicsr_dict.items():
|
||||
if value == k:
|
||||
new_state_dict[key] = v
|
||||
|
||||
elif k.startswith('body.'):
|
||||
bsr_name = k.replace('body.', 'model.1.sub.').lower()
|
||||
bsr_name = bsr_name.replace('rdb', 'RDB')
|
||||
bsr_name = bsr_name.replace('.weight', '.0.weight')
|
||||
bsr_name = bsr_name.replace('.bias', '.0.bias')
|
||||
new_state_dict[bsr_name] = v
|
||||
else:
|
||||
new_state_dict[k] = v
|
||||
return new_state_dict
|
||||
57
toolkit/extension.py
Normal file
57
toolkit/extension.py
Normal file
@@ -0,0 +1,57 @@
|
||||
import os
|
||||
import importlib
|
||||
import pkgutil
|
||||
from typing import List
|
||||
|
||||
from toolkit.paths import TOOLKIT_ROOT
|
||||
|
||||
|
||||
class Extension(object):
|
||||
"""Base class for extensions.
|
||||
|
||||
Extensions are registered with the ExtensionManager, which is
|
||||
responsible for calling the extension's load() and unload()
|
||||
methods at the appropriate times.
|
||||
|
||||
"""
|
||||
|
||||
name: str = None
|
||||
uid: str = None
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# extend in subclass
|
||||
pass
|
||||
|
||||
|
||||
def get_all_extensions() -> List[Extension]:
|
||||
extension_folders = ['extensions', 'extensions_built_in']
|
||||
|
||||
# This will hold the classes from all extension modules
|
||||
all_extension_classes: List[Extension] = []
|
||||
|
||||
# Iterate over all directories (i.e., packages) in the "extensions" directory
|
||||
for sub_dir in extension_folders:
|
||||
extensions_dir = os.path.join(TOOLKIT_ROOT, sub_dir)
|
||||
for (_, name, _) in pkgutil.iter_modules([extensions_dir]):
|
||||
try:
|
||||
# Import the module
|
||||
module = importlib.import_module(f"{sub_dir}.{name}")
|
||||
# Get the value of the AI_TOOLKIT_EXTENSIONS variable
|
||||
extensions = getattr(module, "AI_TOOLKIT_EXTENSIONS", None)
|
||||
# Check if the value is a list
|
||||
if isinstance(extensions, list):
|
||||
# Iterate over the list and add the classes to the main list
|
||||
all_extension_classes.extend(extensions)
|
||||
except ImportError as e:
|
||||
print(f"Failed to import the {name} module. Error: {str(e)}")
|
||||
|
||||
return all_extension_classes
|
||||
|
||||
|
||||
def get_all_extensions_process_dict():
|
||||
all_extensions = get_all_extensions()
|
||||
process_dict = {}
|
||||
for extension in all_extensions:
|
||||
process_dict[extension.uid] = extension.get_process()
|
||||
return process_dict
|
||||
@@ -1,7 +1,12 @@
|
||||
from typing import Union, OrderedDict
|
||||
|
||||
from toolkit.config import get_config
|
||||
|
||||
|
||||
def get_job(config_path, name=None):
|
||||
def get_job(
|
||||
config_path: Union[str, dict, OrderedDict],
|
||||
name=None
|
||||
):
|
||||
config = get_config(config_path, name)
|
||||
if not config['job']:
|
||||
raise ValueError('config file is invalid. Missing "job" key')
|
||||
@@ -16,9 +21,24 @@ def get_job(config_path, name=None):
|
||||
if job == 'mod':
|
||||
from jobs import ModJob
|
||||
return ModJob(config)
|
||||
if job == 'generate':
|
||||
from jobs import GenerateJob
|
||||
return GenerateJob(config)
|
||||
if job == 'extension':
|
||||
from jobs import ExtensionJob
|
||||
return ExtensionJob(config)
|
||||
|
||||
# elif job == 'train':
|
||||
# from jobs import TrainJob
|
||||
# return TrainJob(config)
|
||||
else:
|
||||
raise ValueError(f'Unknown job type {job}')
|
||||
|
||||
|
||||
def run_job(
|
||||
config: Union[str, dict, OrderedDict],
|
||||
name=None
|
||||
):
|
||||
job = get_job(config, name)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
|
||||
@@ -892,6 +892,9 @@ def convert_ldm_clip_checkpoint_v1(checkpoint):
|
||||
for key in keys:
|
||||
if key.startswith("cond_stage_model.transformer"):
|
||||
text_model_dict[key[len("cond_stage_model.transformer."):]] = checkpoint[key]
|
||||
# support checkpoint without position_ids (invalid checkpoint)
|
||||
if "text_model.embeddings.position_ids" not in text_model_dict:
|
||||
text_model_dict["text_model.embeddings.position_ids"] = torch.arange(77).unsqueeze(0) # 77 is the max length of the text
|
||||
return text_model_dict
|
||||
|
||||
|
||||
@@ -1257,6 +1260,10 @@ def load_models_from_stable_diffusion_checkpoint(v2, ckpt_path, device="cpu", dt
|
||||
text_model = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14").to(device)
|
||||
logging.set_verbosity_warning()
|
||||
|
||||
# latest transformers doesnt have position ids. Do we remove it?
|
||||
if "text_model.embeddings.position_ids" not in text_model.state_dict():
|
||||
del converted_text_encoder_checkpoint["text_model.embeddings.position_ids"]
|
||||
|
||||
info = text_model.load_state_dict(converted_text_encoder_checkpoint)
|
||||
print("loading text encoder:", info)
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
class ReductionKernel(nn.Module):
|
||||
@@ -29,3 +30,15 @@ class ReductionKernel(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return nn.functional.conv2d(x, self.kernel, stride=self.kernel_size, padding=0, groups=1)
|
||||
|
||||
|
||||
class CheckpointGradients(nn.Module):
|
||||
def __init__(self, is_gradient_checkpointing=True):
|
||||
super(CheckpointGradients, self).__init__()
|
||||
self.is_gradient_checkpointing = is_gradient_checkpointing
|
||||
|
||||
def forward(self, module, *args, num_chunks=1):
|
||||
if self.is_gradient_checkpointing:
|
||||
return checkpoint(module, *args, num_chunks=self.num_chunks)
|
||||
else:
|
||||
return module(*args)
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from typing import List, Optional, Dict, Type, Union
|
||||
|
||||
@@ -9,7 +11,174 @@ from .paths import SD_SCRIPTS_ROOT
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from networks.lora import LoRANetwork, LoRAModule, get_block_index
|
||||
from networks.lora import LoRANetwork, get_block_index
|
||||
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
|
||||
|
||||
|
||||
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
|
||||
|
||||
# if limit_rank:
|
||||
# self.lora_dim = min(lora_dim, in_dim, out_dim)
|
||||
# if self.lora_dim != lora_dim:
|
||||
# print(f"{lora_name} dim (rank) is changed to: {self.lora_dim}")
|
||||
# else:
|
||||
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: Union[float, List[float]] = multiplier
|
||||
self.org_module = org_module # remove in applying
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
def apply_to(self):
|
||||
self.org_forward = self.org_module.forward
|
||||
self.org_module.forward = self.forward
|
||||
del self.org_module
|
||||
|
||||
# this allows us to set different multipliers on a per item in a batch basis
|
||||
# allowing us to run positive and negative weights in the same batch
|
||||
# really only useful for slider training for now
|
||||
def get_multiplier(self, lora_up):
|
||||
batch_size = lora_up.size(0)
|
||||
# batch will have all negative prompts first and positive prompts second
|
||||
# our multiplier list is for a prompt pair. So we need to repeat it for positive and negative prompts
|
||||
# if there is more than our multiplier, it is liekly a batch size increase, so we need to
|
||||
# interleve the multipliers
|
||||
if isinstance(self.multiplier, list):
|
||||
if len(self.multiplier) == 0:
|
||||
# single item, just return it
|
||||
return self.multiplier[0]
|
||||
elif len(self.multiplier) == batch_size:
|
||||
# not doing CFG
|
||||
multiplier_tensor = torch.tensor(self.multiplier).to(lora_up.device, dtype=lora_up.dtype)
|
||||
else:
|
||||
|
||||
# we have a list of multipliers, so we need to get the multiplier for this batch
|
||||
multiplier_tensor = torch.tensor(self.multiplier * 2).to(lora_up.device, dtype=lora_up.dtype)
|
||||
# should be 1 for if total batch size was 1
|
||||
num_interleaves = (batch_size // 2) // len(self.multiplier)
|
||||
multiplier_tensor = multiplier_tensor.repeat_interleave(num_interleaves)
|
||||
|
||||
# match lora_up rank
|
||||
if len(lora_up.size()) == 2:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1)
|
||||
elif len(lora_up.size()) == 3:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1, 1)
|
||||
elif len(lora_up.size()) == 4:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1, 1, 1)
|
||||
return multiplier_tensor
|
||||
|
||||
else:
|
||||
return self.multiplier
|
||||
|
||||
def _call_forward(self, x):
|
||||
# module dropout
|
||||
if self.module_dropout is not None and self.training:
|
||||
if torch.rand(1) < self.module_dropout:
|
||||
return 0.0 # added to original forward
|
||||
|
||||
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)
|
||||
|
||||
multiplier = self.get_multiplier(lx)
|
||||
|
||||
return lx * multiplier * scale
|
||||
|
||||
def create_custom_forward(self):
|
||||
def custom_forward(*inputs):
|
||||
return self._call_forward(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
def forward(self, x):
|
||||
org_forwarded = self.org_forward(x)
|
||||
# TODO this just loses the grad. Not sure why. Probably why no one else is doing it either
|
||||
# if torch.is_grad_enabled() and self.is_checkpointing and self.training:
|
||||
# lora_output = checkpoint(
|
||||
# self.create_custom_forward(),
|
||||
# x,
|
||||
# )
|
||||
# else:
|
||||
# lora_output = self._call_forward(x)
|
||||
|
||||
lora_output = self._call_forward(x)
|
||||
|
||||
return org_forwarded + lora_output
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.is_checkpointing = True
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.is_checkpointing = False
|
||||
|
||||
|
||||
class LoRASpecialNetwork(LoRANetwork):
|
||||
@@ -70,6 +239,7 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
if modules_dim is not None:
|
||||
print(f"create LoRA network from weights")
|
||||
@@ -236,11 +406,11 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
torch.save(state_dict, file)
|
||||
|
||||
@property
|
||||
def multiplier(self):
|
||||
def multiplier(self) -> Union[float, List[float]]:
|
||||
return self._multiplier
|
||||
|
||||
@multiplier.setter
|
||||
def multiplier(self, value):
|
||||
def multiplier(self, value: Union[float, List[float]]):
|
||||
self._multiplier = value
|
||||
self._update_lora_multiplier()
|
||||
|
||||
@@ -261,6 +431,8 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.multiplier = 0
|
||||
|
||||
# called when the context manager is entered
|
||||
# ie: with network:
|
||||
def __enter__(self):
|
||||
self.is_active = True
|
||||
self._update_lora_multiplier()
|
||||
@@ -278,3 +450,29 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
loras += self.text_encoder_loras
|
||||
for lora in loras:
|
||||
lora.to(device, dtype)
|
||||
|
||||
def _update_checkpointing(self):
|
||||
if self.is_checkpointing:
|
||||
if hasattr(self, 'unet_loras'):
|
||||
for lora in self.unet_loras:
|
||||
lora.enable_gradient_checkpointing()
|
||||
if hasattr(self, 'text_encoder_loras'):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.enable_gradient_checkpointing()
|
||||
else:
|
||||
if hasattr(self, 'unet_loras'):
|
||||
for lora in self.unet_loras:
|
||||
lora.disable_gradient_checkpointing()
|
||||
if hasattr(self, 'text_encoder_loras'):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.disable_gradient_checkpointing()
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
# not supported
|
||||
self.is_checkpointing = True
|
||||
self._update_checkpointing()
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
# not supported
|
||||
self.is_checkpointing = False
|
||||
self._update_checkpointing()
|
||||
|
||||
130
toolkit/model_util_sdxl.py
Normal file
130
toolkit/model_util_sdxl.py
Normal file
@@ -0,0 +1,130 @@
|
||||
import torch
|
||||
from diffusers import AutoencoderKL
|
||||
from safetensors.torch import load_file
|
||||
from transformers import CLIPTextModelWithProjection, CLIPTextConfig, CLIPTextModel
|
||||
|
||||
from library import model_util, sdxl_original_unet
|
||||
from library.sdxl_model_util import convert_sdxl_text_encoder_2_checkpoint
|
||||
|
||||
|
||||
def load_models_from_sdxl_checkpoint(model_version, ckpt_path, map_location):
|
||||
# model_version is reserved for future use
|
||||
|
||||
# Load the state dict
|
||||
if model_util.is_safetensors(ckpt_path):
|
||||
checkpoint = None
|
||||
state_dict = load_file(ckpt_path, device=map_location)
|
||||
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
|
||||
print("building U-Net")
|
||||
unet = sdxl_original_unet.SdxlUNet2DConditionModel()
|
||||
|
||||
print("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 = unet.load_state_dict(unet_sd)
|
||||
print("U-Net: ", info)
|
||||
del unet_sd
|
||||
|
||||
# Text Encoders
|
||||
print("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",
|
||||
)
|
||||
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",
|
||||
)
|
||||
text_model2 = CLIPTextModelWithProjection(text_model2_cfg)
|
||||
|
||||
print("loading text encoders from checkpoint")
|
||||
te1_sd = {}
|
||||
te2_sd = {}
|
||||
for k in list(state_dict.keys()):
|
||||
if k.endswith("text_model.embeddings.position_ids"):
|
||||
# skip position_ids
|
||||
state_dict.pop(k)
|
||||
elif 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)
|
||||
|
||||
|
||||
|
||||
info1 = text_model1.load_state_dict(te1_sd)
|
||||
print("text encoder 1:", info1)
|
||||
|
||||
converted_sd, logit_scale = convert_sdxl_text_encoder_2_checkpoint(te2_sd, max_length=77)
|
||||
# remove text_model.embeddings.position_ids"
|
||||
converted_sd.pop("text_model.embeddings.position_ids")
|
||||
info2 = text_model2.load_state_dict(converted_sd)
|
||||
print("text encoder 2:", info2)
|
||||
|
||||
# prepare vae
|
||||
print("building VAE")
|
||||
vae_config = model_util.create_vae_diffusers_config()
|
||||
vae = AutoencoderKL(**vae_config) # .to(device)
|
||||
|
||||
print("loading VAE from checkpoint")
|
||||
converted_vae_checkpoint = model_util.convert_ldm_vae_checkpoint(state_dict, vae_config)
|
||||
info = vae.load_state_dict(converted_vae_checkpoint)
|
||||
print("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
|
||||
645
toolkit/models/RRDB.py
Normal file
645
toolkit/models/RRDB.py
Normal file
@@ -0,0 +1,645 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import functools
|
||||
import math
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from . import block as B
|
||||
|
||||
esrgan_safetensors_keys = ['model.0.weight', 'model.0.bias', 'model.1.sub.0.RDB1.conv1.0.weight',
|
||||
'model.1.sub.0.RDB1.conv1.0.bias', 'model.1.sub.0.RDB1.conv2.0.weight',
|
||||
'model.1.sub.0.RDB1.conv2.0.bias', 'model.1.sub.0.RDB1.conv3.0.weight',
|
||||
'model.1.sub.0.RDB1.conv3.0.bias', 'model.1.sub.0.RDB1.conv4.0.weight',
|
||||
'model.1.sub.0.RDB1.conv4.0.bias', 'model.1.sub.0.RDB1.conv5.0.weight',
|
||||
'model.1.sub.0.RDB1.conv5.0.bias', 'model.1.sub.0.RDB2.conv1.0.weight',
|
||||
'model.1.sub.0.RDB2.conv1.0.bias', 'model.1.sub.0.RDB2.conv2.0.weight',
|
||||
'model.1.sub.0.RDB2.conv2.0.bias', 'model.1.sub.0.RDB2.conv3.0.weight',
|
||||
'model.1.sub.0.RDB2.conv3.0.bias', 'model.1.sub.0.RDB2.conv4.0.weight',
|
||||
'model.1.sub.0.RDB2.conv4.0.bias', 'model.1.sub.0.RDB2.conv5.0.weight',
|
||||
'model.1.sub.0.RDB2.conv5.0.bias', 'model.1.sub.0.RDB3.conv1.0.weight',
|
||||
'model.1.sub.0.RDB3.conv1.0.bias', 'model.1.sub.0.RDB3.conv2.0.weight',
|
||||
'model.1.sub.0.RDB3.conv2.0.bias', 'model.1.sub.0.RDB3.conv3.0.weight',
|
||||
'model.1.sub.0.RDB3.conv3.0.bias', 'model.1.sub.0.RDB3.conv4.0.weight',
|
||||
'model.1.sub.0.RDB3.conv4.0.bias', 'model.1.sub.0.RDB3.conv5.0.weight',
|
||||
'model.1.sub.0.RDB3.conv5.0.bias', 'model.1.sub.1.RDB1.conv1.0.weight',
|
||||
'model.1.sub.1.RDB1.conv1.0.bias', 'model.1.sub.1.RDB1.conv2.0.weight',
|
||||
'model.1.sub.1.RDB1.conv2.0.bias', 'model.1.sub.1.RDB1.conv3.0.weight',
|
||||
'model.1.sub.1.RDB1.conv3.0.bias', 'model.1.sub.1.RDB1.conv4.0.weight',
|
||||
'model.1.sub.1.RDB1.conv4.0.bias', 'model.1.sub.1.RDB1.conv5.0.weight',
|
||||
'model.1.sub.1.RDB1.conv5.0.bias', 'model.1.sub.1.RDB2.conv1.0.weight',
|
||||
'model.1.sub.1.RDB2.conv1.0.bias', 'model.1.sub.1.RDB2.conv2.0.weight',
|
||||
'model.1.sub.1.RDB2.conv2.0.bias', 'model.1.sub.1.RDB2.conv3.0.weight',
|
||||
'model.1.sub.1.RDB2.conv3.0.bias', 'model.1.sub.1.RDB2.conv4.0.weight',
|
||||
'model.1.sub.1.RDB2.conv4.0.bias', 'model.1.sub.1.RDB2.conv5.0.weight',
|
||||
'model.1.sub.1.RDB2.conv5.0.bias', 'model.1.sub.1.RDB3.conv1.0.weight',
|
||||
'model.1.sub.1.RDB3.conv1.0.bias', 'model.1.sub.1.RDB3.conv2.0.weight',
|
||||
'model.1.sub.1.RDB3.conv2.0.bias', 'model.1.sub.1.RDB3.conv3.0.weight',
|
||||
'model.1.sub.1.RDB3.conv3.0.bias', 'model.1.sub.1.RDB3.conv4.0.weight',
|
||||
'model.1.sub.1.RDB3.conv4.0.bias', 'model.1.sub.1.RDB3.conv5.0.weight',
|
||||
'model.1.sub.1.RDB3.conv5.0.bias', 'model.1.sub.2.RDB1.conv1.0.weight',
|
||||
'model.1.sub.2.RDB1.conv1.0.bias', 'model.1.sub.2.RDB1.conv2.0.weight',
|
||||
'model.1.sub.2.RDB1.conv2.0.bias', 'model.1.sub.2.RDB1.conv3.0.weight',
|
||||
'model.1.sub.2.RDB1.conv3.0.bias', 'model.1.sub.2.RDB1.conv4.0.weight',
|
||||
'model.1.sub.2.RDB1.conv4.0.bias', 'model.1.sub.2.RDB1.conv5.0.weight',
|
||||
'model.1.sub.2.RDB1.conv5.0.bias', 'model.1.sub.2.RDB2.conv1.0.weight',
|
||||
'model.1.sub.2.RDB2.conv1.0.bias', 'model.1.sub.2.RDB2.conv2.0.weight',
|
||||
'model.1.sub.2.RDB2.conv2.0.bias', 'model.1.sub.2.RDB2.conv3.0.weight',
|
||||
'model.1.sub.2.RDB2.conv3.0.bias', 'model.1.sub.2.RDB2.conv4.0.weight',
|
||||
'model.1.sub.2.RDB2.conv4.0.bias', 'model.1.sub.2.RDB2.conv5.0.weight',
|
||||
'model.1.sub.2.RDB2.conv5.0.bias', 'model.1.sub.2.RDB3.conv1.0.weight',
|
||||
'model.1.sub.2.RDB3.conv1.0.bias', 'model.1.sub.2.RDB3.conv2.0.weight',
|
||||
'model.1.sub.2.RDB3.conv2.0.bias', 'model.1.sub.2.RDB3.conv3.0.weight',
|
||||
'model.1.sub.2.RDB3.conv3.0.bias', 'model.1.sub.2.RDB3.conv4.0.weight',
|
||||
'model.1.sub.2.RDB3.conv4.0.bias', 'model.1.sub.2.RDB3.conv5.0.weight',
|
||||
'model.1.sub.2.RDB3.conv5.0.bias', 'model.1.sub.3.RDB1.conv1.0.weight',
|
||||
'model.1.sub.3.RDB1.conv1.0.bias', 'model.1.sub.3.RDB1.conv2.0.weight',
|
||||
'model.1.sub.3.RDB1.conv2.0.bias', 'model.1.sub.3.RDB1.conv3.0.weight',
|
||||
'model.1.sub.3.RDB1.conv3.0.bias', 'model.1.sub.3.RDB1.conv4.0.weight',
|
||||
'model.1.sub.3.RDB1.conv4.0.bias', 'model.1.sub.3.RDB1.conv5.0.weight',
|
||||
'model.1.sub.3.RDB1.conv5.0.bias', 'model.1.sub.3.RDB2.conv1.0.weight',
|
||||
'model.1.sub.3.RDB2.conv1.0.bias', 'model.1.sub.3.RDB2.conv2.0.weight',
|
||||
'model.1.sub.3.RDB2.conv2.0.bias', 'model.1.sub.3.RDB2.conv3.0.weight',
|
||||
'model.1.sub.3.RDB2.conv3.0.bias', 'model.1.sub.3.RDB2.conv4.0.weight',
|
||||
'model.1.sub.3.RDB2.conv4.0.bias', 'model.1.sub.3.RDB2.conv5.0.weight',
|
||||
'model.1.sub.3.RDB2.conv5.0.bias', 'model.1.sub.3.RDB3.conv1.0.weight',
|
||||
'model.1.sub.3.RDB3.conv1.0.bias', 'model.1.sub.3.RDB3.conv2.0.weight',
|
||||
'model.1.sub.3.RDB3.conv2.0.bias', 'model.1.sub.3.RDB3.conv3.0.weight',
|
||||
'model.1.sub.3.RDB3.conv3.0.bias', 'model.1.sub.3.RDB3.conv4.0.weight',
|
||||
'model.1.sub.3.RDB3.conv4.0.bias', 'model.1.sub.3.RDB3.conv5.0.weight',
|
||||
'model.1.sub.3.RDB3.conv5.0.bias', 'model.1.sub.4.RDB1.conv1.0.weight',
|
||||
'model.1.sub.4.RDB1.conv1.0.bias', 'model.1.sub.4.RDB1.conv2.0.weight',
|
||||
'model.1.sub.4.RDB1.conv2.0.bias', 'model.1.sub.4.RDB1.conv3.0.weight',
|
||||
'model.1.sub.4.RDB1.conv3.0.bias', 'model.1.sub.4.RDB1.conv4.0.weight',
|
||||
'model.1.sub.4.RDB1.conv4.0.bias', 'model.1.sub.4.RDB1.conv5.0.weight',
|
||||
'model.1.sub.4.RDB1.conv5.0.bias', 'model.1.sub.4.RDB2.conv1.0.weight',
|
||||
'model.1.sub.4.RDB2.conv1.0.bias', 'model.1.sub.4.RDB2.conv2.0.weight',
|
||||
'model.1.sub.4.RDB2.conv2.0.bias', 'model.1.sub.4.RDB2.conv3.0.weight',
|
||||
'model.1.sub.4.RDB2.conv3.0.bias', 'model.1.sub.4.RDB2.conv4.0.weight',
|
||||
'model.1.sub.4.RDB2.conv4.0.bias', 'model.1.sub.4.RDB2.conv5.0.weight',
|
||||
'model.1.sub.4.RDB2.conv5.0.bias', 'model.1.sub.4.RDB3.conv1.0.weight',
|
||||
'model.1.sub.4.RDB3.conv1.0.bias', 'model.1.sub.4.RDB3.conv2.0.weight',
|
||||
'model.1.sub.4.RDB3.conv2.0.bias', 'model.1.sub.4.RDB3.conv3.0.weight',
|
||||
'model.1.sub.4.RDB3.conv3.0.bias', 'model.1.sub.4.RDB3.conv4.0.weight',
|
||||
'model.1.sub.4.RDB3.conv4.0.bias', 'model.1.sub.4.RDB3.conv5.0.weight',
|
||||
'model.1.sub.4.RDB3.conv5.0.bias', 'model.1.sub.5.RDB1.conv1.0.weight',
|
||||
'model.1.sub.5.RDB1.conv1.0.bias', 'model.1.sub.5.RDB1.conv2.0.weight',
|
||||
'model.1.sub.5.RDB1.conv2.0.bias', 'model.1.sub.5.RDB1.conv3.0.weight',
|
||||
'model.1.sub.5.RDB1.conv3.0.bias', 'model.1.sub.5.RDB1.conv4.0.weight',
|
||||
'model.1.sub.5.RDB1.conv4.0.bias', 'model.1.sub.5.RDB1.conv5.0.weight',
|
||||
'model.1.sub.5.RDB1.conv5.0.bias', 'model.1.sub.5.RDB2.conv1.0.weight',
|
||||
'model.1.sub.5.RDB2.conv1.0.bias', 'model.1.sub.5.RDB2.conv2.0.weight',
|
||||
'model.1.sub.5.RDB2.conv2.0.bias', 'model.1.sub.5.RDB2.conv3.0.weight',
|
||||
'model.1.sub.5.RDB2.conv3.0.bias', 'model.1.sub.5.RDB2.conv4.0.weight',
|
||||
'model.1.sub.5.RDB2.conv4.0.bias', 'model.1.sub.5.RDB2.conv5.0.weight',
|
||||
'model.1.sub.5.RDB2.conv5.0.bias', 'model.1.sub.5.RDB3.conv1.0.weight',
|
||||
'model.1.sub.5.RDB3.conv1.0.bias', 'model.1.sub.5.RDB3.conv2.0.weight',
|
||||
'model.1.sub.5.RDB3.conv2.0.bias', 'model.1.sub.5.RDB3.conv3.0.weight',
|
||||
'model.1.sub.5.RDB3.conv3.0.bias', 'model.1.sub.5.RDB3.conv4.0.weight',
|
||||
'model.1.sub.5.RDB3.conv4.0.bias', 'model.1.sub.5.RDB3.conv5.0.weight',
|
||||
'model.1.sub.5.RDB3.conv5.0.bias', 'model.1.sub.6.RDB1.conv1.0.weight',
|
||||
'model.1.sub.6.RDB1.conv1.0.bias', 'model.1.sub.6.RDB1.conv2.0.weight',
|
||||
'model.1.sub.6.RDB1.conv2.0.bias', 'model.1.sub.6.RDB1.conv3.0.weight',
|
||||
'model.1.sub.6.RDB1.conv3.0.bias', 'model.1.sub.6.RDB1.conv4.0.weight',
|
||||
'model.1.sub.6.RDB1.conv4.0.bias', 'model.1.sub.6.RDB1.conv5.0.weight',
|
||||
'model.1.sub.6.RDB1.conv5.0.bias', 'model.1.sub.6.RDB2.conv1.0.weight',
|
||||
'model.1.sub.6.RDB2.conv1.0.bias', 'model.1.sub.6.RDB2.conv2.0.weight',
|
||||
'model.1.sub.6.RDB2.conv2.0.bias', 'model.1.sub.6.RDB2.conv3.0.weight',
|
||||
'model.1.sub.6.RDB2.conv3.0.bias', 'model.1.sub.6.RDB2.conv4.0.weight',
|
||||
'model.1.sub.6.RDB2.conv4.0.bias', 'model.1.sub.6.RDB2.conv5.0.weight',
|
||||
'model.1.sub.6.RDB2.conv5.0.bias', 'model.1.sub.6.RDB3.conv1.0.weight',
|
||||
'model.1.sub.6.RDB3.conv1.0.bias', 'model.1.sub.6.RDB3.conv2.0.weight',
|
||||
'model.1.sub.6.RDB3.conv2.0.bias', 'model.1.sub.6.RDB3.conv3.0.weight',
|
||||
'model.1.sub.6.RDB3.conv3.0.bias', 'model.1.sub.6.RDB3.conv4.0.weight',
|
||||
'model.1.sub.6.RDB3.conv4.0.bias', 'model.1.sub.6.RDB3.conv5.0.weight',
|
||||
'model.1.sub.6.RDB3.conv5.0.bias', 'model.1.sub.7.RDB1.conv1.0.weight',
|
||||
'model.1.sub.7.RDB1.conv1.0.bias', 'model.1.sub.7.RDB1.conv2.0.weight',
|
||||
'model.1.sub.7.RDB1.conv2.0.bias', 'model.1.sub.7.RDB1.conv3.0.weight',
|
||||
'model.1.sub.7.RDB1.conv3.0.bias', 'model.1.sub.7.RDB1.conv4.0.weight',
|
||||
'model.1.sub.7.RDB1.conv4.0.bias', 'model.1.sub.7.RDB1.conv5.0.weight',
|
||||
'model.1.sub.7.RDB1.conv5.0.bias', 'model.1.sub.7.RDB2.conv1.0.weight',
|
||||
'model.1.sub.7.RDB2.conv1.0.bias', 'model.1.sub.7.RDB2.conv2.0.weight',
|
||||
'model.1.sub.7.RDB2.conv2.0.bias', 'model.1.sub.7.RDB2.conv3.0.weight',
|
||||
'model.1.sub.7.RDB2.conv3.0.bias', 'model.1.sub.7.RDB2.conv4.0.weight',
|
||||
'model.1.sub.7.RDB2.conv4.0.bias', 'model.1.sub.7.RDB2.conv5.0.weight',
|
||||
'model.1.sub.7.RDB2.conv5.0.bias', 'model.1.sub.7.RDB3.conv1.0.weight',
|
||||
'model.1.sub.7.RDB3.conv1.0.bias', 'model.1.sub.7.RDB3.conv2.0.weight',
|
||||
'model.1.sub.7.RDB3.conv2.0.bias', 'model.1.sub.7.RDB3.conv3.0.weight',
|
||||
'model.1.sub.7.RDB3.conv3.0.bias', 'model.1.sub.7.RDB3.conv4.0.weight',
|
||||
'model.1.sub.7.RDB3.conv4.0.bias', 'model.1.sub.7.RDB3.conv5.0.weight',
|
||||
'model.1.sub.7.RDB3.conv5.0.bias', 'model.1.sub.8.RDB1.conv1.0.weight',
|
||||
'model.1.sub.8.RDB1.conv1.0.bias', 'model.1.sub.8.RDB1.conv2.0.weight',
|
||||
'model.1.sub.8.RDB1.conv2.0.bias', 'model.1.sub.8.RDB1.conv3.0.weight',
|
||||
'model.1.sub.8.RDB1.conv3.0.bias', 'model.1.sub.8.RDB1.conv4.0.weight',
|
||||
'model.1.sub.8.RDB1.conv4.0.bias', 'model.1.sub.8.RDB1.conv5.0.weight',
|
||||
'model.1.sub.8.RDB1.conv5.0.bias', 'model.1.sub.8.RDB2.conv1.0.weight',
|
||||
'model.1.sub.8.RDB2.conv1.0.bias', 'model.1.sub.8.RDB2.conv2.0.weight',
|
||||
'model.1.sub.8.RDB2.conv2.0.bias', 'model.1.sub.8.RDB2.conv3.0.weight',
|
||||
'model.1.sub.8.RDB2.conv3.0.bias', 'model.1.sub.8.RDB2.conv4.0.weight',
|
||||
'model.1.sub.8.RDB2.conv4.0.bias', 'model.1.sub.8.RDB2.conv5.0.weight',
|
||||
'model.1.sub.8.RDB2.conv5.0.bias', 'model.1.sub.8.RDB3.conv1.0.weight',
|
||||
'model.1.sub.8.RDB3.conv1.0.bias', 'model.1.sub.8.RDB3.conv2.0.weight',
|
||||
'model.1.sub.8.RDB3.conv2.0.bias', 'model.1.sub.8.RDB3.conv3.0.weight',
|
||||
'model.1.sub.8.RDB3.conv3.0.bias', 'model.1.sub.8.RDB3.conv4.0.weight',
|
||||
'model.1.sub.8.RDB3.conv4.0.bias', 'model.1.sub.8.RDB3.conv5.0.weight',
|
||||
'model.1.sub.8.RDB3.conv5.0.bias', 'model.1.sub.9.RDB1.conv1.0.weight',
|
||||
'model.1.sub.9.RDB1.conv1.0.bias', 'model.1.sub.9.RDB1.conv2.0.weight',
|
||||
'model.1.sub.9.RDB1.conv2.0.bias', 'model.1.sub.9.RDB1.conv3.0.weight',
|
||||
'model.1.sub.9.RDB1.conv3.0.bias', 'model.1.sub.9.RDB1.conv4.0.weight',
|
||||
'model.1.sub.9.RDB1.conv4.0.bias', 'model.1.sub.9.RDB1.conv5.0.weight',
|
||||
'model.1.sub.9.RDB1.conv5.0.bias', 'model.1.sub.9.RDB2.conv1.0.weight',
|
||||
'model.1.sub.9.RDB2.conv1.0.bias', 'model.1.sub.9.RDB2.conv2.0.weight',
|
||||
'model.1.sub.9.RDB2.conv2.0.bias', 'model.1.sub.9.RDB2.conv3.0.weight',
|
||||
'model.1.sub.9.RDB2.conv3.0.bias', 'model.1.sub.9.RDB2.conv4.0.weight',
|
||||
'model.1.sub.9.RDB2.conv4.0.bias', 'model.1.sub.9.RDB2.conv5.0.weight',
|
||||
'model.1.sub.9.RDB2.conv5.0.bias', 'model.1.sub.9.RDB3.conv1.0.weight',
|
||||
'model.1.sub.9.RDB3.conv1.0.bias', 'model.1.sub.9.RDB3.conv2.0.weight',
|
||||
'model.1.sub.9.RDB3.conv2.0.bias', 'model.1.sub.9.RDB3.conv3.0.weight',
|
||||
'model.1.sub.9.RDB3.conv3.0.bias', 'model.1.sub.9.RDB3.conv4.0.weight',
|
||||
'model.1.sub.9.RDB3.conv4.0.bias', 'model.1.sub.9.RDB3.conv5.0.weight',
|
||||
'model.1.sub.9.RDB3.conv5.0.bias', 'model.1.sub.10.RDB1.conv1.0.weight',
|
||||
'model.1.sub.10.RDB1.conv1.0.bias', 'model.1.sub.10.RDB1.conv2.0.weight',
|
||||
'model.1.sub.10.RDB1.conv2.0.bias', 'model.1.sub.10.RDB1.conv3.0.weight',
|
||||
'model.1.sub.10.RDB1.conv3.0.bias', 'model.1.sub.10.RDB1.conv4.0.weight',
|
||||
'model.1.sub.10.RDB1.conv4.0.bias', 'model.1.sub.10.RDB1.conv5.0.weight',
|
||||
'model.1.sub.10.RDB1.conv5.0.bias', 'model.1.sub.10.RDB2.conv1.0.weight',
|
||||
'model.1.sub.10.RDB2.conv1.0.bias', 'model.1.sub.10.RDB2.conv2.0.weight',
|
||||
'model.1.sub.10.RDB2.conv2.0.bias', 'model.1.sub.10.RDB2.conv3.0.weight',
|
||||
'model.1.sub.10.RDB2.conv3.0.bias', 'model.1.sub.10.RDB2.conv4.0.weight',
|
||||
'model.1.sub.10.RDB2.conv4.0.bias', 'model.1.sub.10.RDB2.conv5.0.weight',
|
||||
'model.1.sub.10.RDB2.conv5.0.bias', 'model.1.sub.10.RDB3.conv1.0.weight',
|
||||
'model.1.sub.10.RDB3.conv1.0.bias', 'model.1.sub.10.RDB3.conv2.0.weight',
|
||||
'model.1.sub.10.RDB3.conv2.0.bias', 'model.1.sub.10.RDB3.conv3.0.weight',
|
||||
'model.1.sub.10.RDB3.conv3.0.bias', 'model.1.sub.10.RDB3.conv4.0.weight',
|
||||
'model.1.sub.10.RDB3.conv4.0.bias', 'model.1.sub.10.RDB3.conv5.0.weight',
|
||||
'model.1.sub.10.RDB3.conv5.0.bias', 'model.1.sub.11.RDB1.conv1.0.weight',
|
||||
'model.1.sub.11.RDB1.conv1.0.bias', 'model.1.sub.11.RDB1.conv2.0.weight',
|
||||
'model.1.sub.11.RDB1.conv2.0.bias', 'model.1.sub.11.RDB1.conv3.0.weight',
|
||||
'model.1.sub.11.RDB1.conv3.0.bias', 'model.1.sub.11.RDB1.conv4.0.weight',
|
||||
'model.1.sub.11.RDB1.conv4.0.bias', 'model.1.sub.11.RDB1.conv5.0.weight',
|
||||
'model.1.sub.11.RDB1.conv5.0.bias', 'model.1.sub.11.RDB2.conv1.0.weight',
|
||||
'model.1.sub.11.RDB2.conv1.0.bias', 'model.1.sub.11.RDB2.conv2.0.weight',
|
||||
'model.1.sub.11.RDB2.conv2.0.bias', 'model.1.sub.11.RDB2.conv3.0.weight',
|
||||
'model.1.sub.11.RDB2.conv3.0.bias', 'model.1.sub.11.RDB2.conv4.0.weight',
|
||||
'model.1.sub.11.RDB2.conv4.0.bias', 'model.1.sub.11.RDB2.conv5.0.weight',
|
||||
'model.1.sub.11.RDB2.conv5.0.bias', 'model.1.sub.11.RDB3.conv1.0.weight',
|
||||
'model.1.sub.11.RDB3.conv1.0.bias', 'model.1.sub.11.RDB3.conv2.0.weight',
|
||||
'model.1.sub.11.RDB3.conv2.0.bias', 'model.1.sub.11.RDB3.conv3.0.weight',
|
||||
'model.1.sub.11.RDB3.conv3.0.bias', 'model.1.sub.11.RDB3.conv4.0.weight',
|
||||
'model.1.sub.11.RDB3.conv4.0.bias', 'model.1.sub.11.RDB3.conv5.0.weight',
|
||||
'model.1.sub.11.RDB3.conv5.0.bias', 'model.1.sub.12.RDB1.conv1.0.weight',
|
||||
'model.1.sub.12.RDB1.conv1.0.bias', 'model.1.sub.12.RDB1.conv2.0.weight',
|
||||
'model.1.sub.12.RDB1.conv2.0.bias', 'model.1.sub.12.RDB1.conv3.0.weight',
|
||||
'model.1.sub.12.RDB1.conv3.0.bias', 'model.1.sub.12.RDB1.conv4.0.weight',
|
||||
'model.1.sub.12.RDB1.conv4.0.bias', 'model.1.sub.12.RDB1.conv5.0.weight',
|
||||
'model.1.sub.12.RDB1.conv5.0.bias', 'model.1.sub.12.RDB2.conv1.0.weight',
|
||||
'model.1.sub.12.RDB2.conv1.0.bias', 'model.1.sub.12.RDB2.conv2.0.weight',
|
||||
'model.1.sub.12.RDB2.conv2.0.bias', 'model.1.sub.12.RDB2.conv3.0.weight',
|
||||
'model.1.sub.12.RDB2.conv3.0.bias', 'model.1.sub.12.RDB2.conv4.0.weight',
|
||||
'model.1.sub.12.RDB2.conv4.0.bias', 'model.1.sub.12.RDB2.conv5.0.weight',
|
||||
'model.1.sub.12.RDB2.conv5.0.bias', 'model.1.sub.12.RDB3.conv1.0.weight',
|
||||
'model.1.sub.12.RDB3.conv1.0.bias', 'model.1.sub.12.RDB3.conv2.0.weight',
|
||||
'model.1.sub.12.RDB3.conv2.0.bias', 'model.1.sub.12.RDB3.conv3.0.weight',
|
||||
'model.1.sub.12.RDB3.conv3.0.bias', 'model.1.sub.12.RDB3.conv4.0.weight',
|
||||
'model.1.sub.12.RDB3.conv4.0.bias', 'model.1.sub.12.RDB3.conv5.0.weight',
|
||||
'model.1.sub.12.RDB3.conv5.0.bias', 'model.1.sub.13.RDB1.conv1.0.weight',
|
||||
'model.1.sub.13.RDB1.conv1.0.bias', 'model.1.sub.13.RDB1.conv2.0.weight',
|
||||
'model.1.sub.13.RDB1.conv2.0.bias', 'model.1.sub.13.RDB1.conv3.0.weight',
|
||||
'model.1.sub.13.RDB1.conv3.0.bias', 'model.1.sub.13.RDB1.conv4.0.weight',
|
||||
'model.1.sub.13.RDB1.conv4.0.bias', 'model.1.sub.13.RDB1.conv5.0.weight',
|
||||
'model.1.sub.13.RDB1.conv5.0.bias', 'model.1.sub.13.RDB2.conv1.0.weight',
|
||||
'model.1.sub.13.RDB2.conv1.0.bias', 'model.1.sub.13.RDB2.conv2.0.weight',
|
||||
'model.1.sub.13.RDB2.conv2.0.bias', 'model.1.sub.13.RDB2.conv3.0.weight',
|
||||
'model.1.sub.13.RDB2.conv3.0.bias', 'model.1.sub.13.RDB2.conv4.0.weight',
|
||||
'model.1.sub.13.RDB2.conv4.0.bias', 'model.1.sub.13.RDB2.conv5.0.weight',
|
||||
'model.1.sub.13.RDB2.conv5.0.bias', 'model.1.sub.13.RDB3.conv1.0.weight',
|
||||
'model.1.sub.13.RDB3.conv1.0.bias', 'model.1.sub.13.RDB3.conv2.0.weight',
|
||||
'model.1.sub.13.RDB3.conv2.0.bias', 'model.1.sub.13.RDB3.conv3.0.weight',
|
||||
'model.1.sub.13.RDB3.conv3.0.bias', 'model.1.sub.13.RDB3.conv4.0.weight',
|
||||
'model.1.sub.13.RDB3.conv4.0.bias', 'model.1.sub.13.RDB3.conv5.0.weight',
|
||||
'model.1.sub.13.RDB3.conv5.0.bias', 'model.1.sub.14.RDB1.conv1.0.weight',
|
||||
'model.1.sub.14.RDB1.conv1.0.bias', 'model.1.sub.14.RDB1.conv2.0.weight',
|
||||
'model.1.sub.14.RDB1.conv2.0.bias', 'model.1.sub.14.RDB1.conv3.0.weight',
|
||||
'model.1.sub.14.RDB1.conv3.0.bias', 'model.1.sub.14.RDB1.conv4.0.weight',
|
||||
'model.1.sub.14.RDB1.conv4.0.bias', 'model.1.sub.14.RDB1.conv5.0.weight',
|
||||
'model.1.sub.14.RDB1.conv5.0.bias', 'model.1.sub.14.RDB2.conv1.0.weight',
|
||||
'model.1.sub.14.RDB2.conv1.0.bias', 'model.1.sub.14.RDB2.conv2.0.weight',
|
||||
'model.1.sub.14.RDB2.conv2.0.bias', 'model.1.sub.14.RDB2.conv3.0.weight',
|
||||
'model.1.sub.14.RDB2.conv3.0.bias', 'model.1.sub.14.RDB2.conv4.0.weight',
|
||||
'model.1.sub.14.RDB2.conv4.0.bias', 'model.1.sub.14.RDB2.conv5.0.weight',
|
||||
'model.1.sub.14.RDB2.conv5.0.bias', 'model.1.sub.14.RDB3.conv1.0.weight',
|
||||
'model.1.sub.14.RDB3.conv1.0.bias', 'model.1.sub.14.RDB3.conv2.0.weight',
|
||||
'model.1.sub.14.RDB3.conv2.0.bias', 'model.1.sub.14.RDB3.conv3.0.weight',
|
||||
'model.1.sub.14.RDB3.conv3.0.bias', 'model.1.sub.14.RDB3.conv4.0.weight',
|
||||
'model.1.sub.14.RDB3.conv4.0.bias', 'model.1.sub.14.RDB3.conv5.0.weight',
|
||||
'model.1.sub.14.RDB3.conv5.0.bias', 'model.1.sub.15.RDB1.conv1.0.weight',
|
||||
'model.1.sub.15.RDB1.conv1.0.bias', 'model.1.sub.15.RDB1.conv2.0.weight',
|
||||
'model.1.sub.15.RDB1.conv2.0.bias', 'model.1.sub.15.RDB1.conv3.0.weight',
|
||||
'model.1.sub.15.RDB1.conv3.0.bias', 'model.1.sub.15.RDB1.conv4.0.weight',
|
||||
'model.1.sub.15.RDB1.conv4.0.bias', 'model.1.sub.15.RDB1.conv5.0.weight',
|
||||
'model.1.sub.15.RDB1.conv5.0.bias', 'model.1.sub.15.RDB2.conv1.0.weight',
|
||||
'model.1.sub.15.RDB2.conv1.0.bias', 'model.1.sub.15.RDB2.conv2.0.weight',
|
||||
'model.1.sub.15.RDB2.conv2.0.bias', 'model.1.sub.15.RDB2.conv3.0.weight',
|
||||
'model.1.sub.15.RDB2.conv3.0.bias', 'model.1.sub.15.RDB2.conv4.0.weight',
|
||||
'model.1.sub.15.RDB2.conv4.0.bias', 'model.1.sub.15.RDB2.conv5.0.weight',
|
||||
'model.1.sub.15.RDB2.conv5.0.bias', 'model.1.sub.15.RDB3.conv1.0.weight',
|
||||
'model.1.sub.15.RDB3.conv1.0.bias', 'model.1.sub.15.RDB3.conv2.0.weight',
|
||||
'model.1.sub.15.RDB3.conv2.0.bias', 'model.1.sub.15.RDB3.conv3.0.weight',
|
||||
'model.1.sub.15.RDB3.conv3.0.bias', 'model.1.sub.15.RDB3.conv4.0.weight',
|
||||
'model.1.sub.15.RDB3.conv4.0.bias', 'model.1.sub.15.RDB3.conv5.0.weight',
|
||||
'model.1.sub.15.RDB3.conv5.0.bias', 'model.1.sub.16.RDB1.conv1.0.weight',
|
||||
'model.1.sub.16.RDB1.conv1.0.bias', 'model.1.sub.16.RDB1.conv2.0.weight',
|
||||
'model.1.sub.16.RDB1.conv2.0.bias', 'model.1.sub.16.RDB1.conv3.0.weight',
|
||||
'model.1.sub.16.RDB1.conv3.0.bias', 'model.1.sub.16.RDB1.conv4.0.weight',
|
||||
'model.1.sub.16.RDB1.conv4.0.bias', 'model.1.sub.16.RDB1.conv5.0.weight',
|
||||
'model.1.sub.16.RDB1.conv5.0.bias', 'model.1.sub.16.RDB2.conv1.0.weight',
|
||||
'model.1.sub.16.RDB2.conv1.0.bias', 'model.1.sub.16.RDB2.conv2.0.weight',
|
||||
'model.1.sub.16.RDB2.conv2.0.bias', 'model.1.sub.16.RDB2.conv3.0.weight',
|
||||
'model.1.sub.16.RDB2.conv3.0.bias', 'model.1.sub.16.RDB2.conv4.0.weight',
|
||||
'model.1.sub.16.RDB2.conv4.0.bias', 'model.1.sub.16.RDB2.conv5.0.weight',
|
||||
'model.1.sub.16.RDB2.conv5.0.bias', 'model.1.sub.16.RDB3.conv1.0.weight',
|
||||
'model.1.sub.16.RDB3.conv1.0.bias', 'model.1.sub.16.RDB3.conv2.0.weight',
|
||||
'model.1.sub.16.RDB3.conv2.0.bias', 'model.1.sub.16.RDB3.conv3.0.weight',
|
||||
'model.1.sub.16.RDB3.conv3.0.bias', 'model.1.sub.16.RDB3.conv4.0.weight',
|
||||
'model.1.sub.16.RDB3.conv4.0.bias', 'model.1.sub.16.RDB3.conv5.0.weight',
|
||||
'model.1.sub.16.RDB3.conv5.0.bias', 'model.1.sub.17.RDB1.conv1.0.weight',
|
||||
'model.1.sub.17.RDB1.conv1.0.bias', 'model.1.sub.17.RDB1.conv2.0.weight',
|
||||
'model.1.sub.17.RDB1.conv2.0.bias', 'model.1.sub.17.RDB1.conv3.0.weight',
|
||||
'model.1.sub.17.RDB1.conv3.0.bias', 'model.1.sub.17.RDB1.conv4.0.weight',
|
||||
'model.1.sub.17.RDB1.conv4.0.bias', 'model.1.sub.17.RDB1.conv5.0.weight',
|
||||
'model.1.sub.17.RDB1.conv5.0.bias', 'model.1.sub.17.RDB2.conv1.0.weight',
|
||||
'model.1.sub.17.RDB2.conv1.0.bias', 'model.1.sub.17.RDB2.conv2.0.weight',
|
||||
'model.1.sub.17.RDB2.conv2.0.bias', 'model.1.sub.17.RDB2.conv3.0.weight',
|
||||
'model.1.sub.17.RDB2.conv3.0.bias', 'model.1.sub.17.RDB2.conv4.0.weight',
|
||||
'model.1.sub.17.RDB2.conv4.0.bias', 'model.1.sub.17.RDB2.conv5.0.weight',
|
||||
'model.1.sub.17.RDB2.conv5.0.bias', 'model.1.sub.17.RDB3.conv1.0.weight',
|
||||
'model.1.sub.17.RDB3.conv1.0.bias', 'model.1.sub.17.RDB3.conv2.0.weight',
|
||||
'model.1.sub.17.RDB3.conv2.0.bias', 'model.1.sub.17.RDB3.conv3.0.weight',
|
||||
'model.1.sub.17.RDB3.conv3.0.bias', 'model.1.sub.17.RDB3.conv4.0.weight',
|
||||
'model.1.sub.17.RDB3.conv4.0.bias', 'model.1.sub.17.RDB3.conv5.0.weight',
|
||||
'model.1.sub.17.RDB3.conv5.0.bias', 'model.1.sub.18.RDB1.conv1.0.weight',
|
||||
'model.1.sub.18.RDB1.conv1.0.bias', 'model.1.sub.18.RDB1.conv2.0.weight',
|
||||
'model.1.sub.18.RDB1.conv2.0.bias', 'model.1.sub.18.RDB1.conv3.0.weight',
|
||||
'model.1.sub.18.RDB1.conv3.0.bias', 'model.1.sub.18.RDB1.conv4.0.weight',
|
||||
'model.1.sub.18.RDB1.conv4.0.bias', 'model.1.sub.18.RDB1.conv5.0.weight',
|
||||
'model.1.sub.18.RDB1.conv5.0.bias', 'model.1.sub.18.RDB2.conv1.0.weight',
|
||||
'model.1.sub.18.RDB2.conv1.0.bias', 'model.1.sub.18.RDB2.conv2.0.weight',
|
||||
'model.1.sub.18.RDB2.conv2.0.bias', 'model.1.sub.18.RDB2.conv3.0.weight',
|
||||
'model.1.sub.18.RDB2.conv3.0.bias', 'model.1.sub.18.RDB2.conv4.0.weight',
|
||||
'model.1.sub.18.RDB2.conv4.0.bias', 'model.1.sub.18.RDB2.conv5.0.weight',
|
||||
'model.1.sub.18.RDB2.conv5.0.bias', 'model.1.sub.18.RDB3.conv1.0.weight',
|
||||
'model.1.sub.18.RDB3.conv1.0.bias', 'model.1.sub.18.RDB3.conv2.0.weight',
|
||||
'model.1.sub.18.RDB3.conv2.0.bias', 'model.1.sub.18.RDB3.conv3.0.weight',
|
||||
'model.1.sub.18.RDB3.conv3.0.bias', 'model.1.sub.18.RDB3.conv4.0.weight',
|
||||
'model.1.sub.18.RDB3.conv4.0.bias', 'model.1.sub.18.RDB3.conv5.0.weight',
|
||||
'model.1.sub.18.RDB3.conv5.0.bias', 'model.1.sub.19.RDB1.conv1.0.weight',
|
||||
'model.1.sub.19.RDB1.conv1.0.bias', 'model.1.sub.19.RDB1.conv2.0.weight',
|
||||
'model.1.sub.19.RDB1.conv2.0.bias', 'model.1.sub.19.RDB1.conv3.0.weight',
|
||||
'model.1.sub.19.RDB1.conv3.0.bias', 'model.1.sub.19.RDB1.conv4.0.weight',
|
||||
'model.1.sub.19.RDB1.conv4.0.bias', 'model.1.sub.19.RDB1.conv5.0.weight',
|
||||
'model.1.sub.19.RDB1.conv5.0.bias', 'model.1.sub.19.RDB2.conv1.0.weight',
|
||||
'model.1.sub.19.RDB2.conv1.0.bias', 'model.1.sub.19.RDB2.conv2.0.weight',
|
||||
'model.1.sub.19.RDB2.conv2.0.bias', 'model.1.sub.19.RDB2.conv3.0.weight',
|
||||
'model.1.sub.19.RDB2.conv3.0.bias', 'model.1.sub.19.RDB2.conv4.0.weight',
|
||||
'model.1.sub.19.RDB2.conv4.0.bias', 'model.1.sub.19.RDB2.conv5.0.weight',
|
||||
'model.1.sub.19.RDB2.conv5.0.bias', 'model.1.sub.19.RDB3.conv1.0.weight',
|
||||
'model.1.sub.19.RDB3.conv1.0.bias', 'model.1.sub.19.RDB3.conv2.0.weight',
|
||||
'model.1.sub.19.RDB3.conv2.0.bias', 'model.1.sub.19.RDB3.conv3.0.weight',
|
||||
'model.1.sub.19.RDB3.conv3.0.bias', 'model.1.sub.19.RDB3.conv4.0.weight',
|
||||
'model.1.sub.19.RDB3.conv4.0.bias', 'model.1.sub.19.RDB3.conv5.0.weight',
|
||||
'model.1.sub.19.RDB3.conv5.0.bias', 'model.1.sub.20.RDB1.conv1.0.weight',
|
||||
'model.1.sub.20.RDB1.conv1.0.bias', 'model.1.sub.20.RDB1.conv2.0.weight',
|
||||
'model.1.sub.20.RDB1.conv2.0.bias', 'model.1.sub.20.RDB1.conv3.0.weight',
|
||||
'model.1.sub.20.RDB1.conv3.0.bias', 'model.1.sub.20.RDB1.conv4.0.weight',
|
||||
'model.1.sub.20.RDB1.conv4.0.bias', 'model.1.sub.20.RDB1.conv5.0.weight',
|
||||
'model.1.sub.20.RDB1.conv5.0.bias', 'model.1.sub.20.RDB2.conv1.0.weight',
|
||||
'model.1.sub.20.RDB2.conv1.0.bias', 'model.1.sub.20.RDB2.conv2.0.weight',
|
||||
'model.1.sub.20.RDB2.conv2.0.bias', 'model.1.sub.20.RDB2.conv3.0.weight',
|
||||
'model.1.sub.20.RDB2.conv3.0.bias', 'model.1.sub.20.RDB2.conv4.0.weight',
|
||||
'model.1.sub.20.RDB2.conv4.0.bias', 'model.1.sub.20.RDB2.conv5.0.weight',
|
||||
'model.1.sub.20.RDB2.conv5.0.bias', 'model.1.sub.20.RDB3.conv1.0.weight',
|
||||
'model.1.sub.20.RDB3.conv1.0.bias', 'model.1.sub.20.RDB3.conv2.0.weight',
|
||||
'model.1.sub.20.RDB3.conv2.0.bias', 'model.1.sub.20.RDB3.conv3.0.weight',
|
||||
'model.1.sub.20.RDB3.conv3.0.bias', 'model.1.sub.20.RDB3.conv4.0.weight',
|
||||
'model.1.sub.20.RDB3.conv4.0.bias', 'model.1.sub.20.RDB3.conv5.0.weight',
|
||||
'model.1.sub.20.RDB3.conv5.0.bias', 'model.1.sub.21.RDB1.conv1.0.weight',
|
||||
'model.1.sub.21.RDB1.conv1.0.bias', 'model.1.sub.21.RDB1.conv2.0.weight',
|
||||
'model.1.sub.21.RDB1.conv2.0.bias', 'model.1.sub.21.RDB1.conv3.0.weight',
|
||||
'model.1.sub.21.RDB1.conv3.0.bias', 'model.1.sub.21.RDB1.conv4.0.weight',
|
||||
'model.1.sub.21.RDB1.conv4.0.bias', 'model.1.sub.21.RDB1.conv5.0.weight',
|
||||
'model.1.sub.21.RDB1.conv5.0.bias', 'model.1.sub.21.RDB2.conv1.0.weight',
|
||||
'model.1.sub.21.RDB2.conv1.0.bias', 'model.1.sub.21.RDB2.conv2.0.weight',
|
||||
'model.1.sub.21.RDB2.conv2.0.bias', 'model.1.sub.21.RDB2.conv3.0.weight',
|
||||
'model.1.sub.21.RDB2.conv3.0.bias', 'model.1.sub.21.RDB2.conv4.0.weight',
|
||||
'model.1.sub.21.RDB2.conv4.0.bias', 'model.1.sub.21.RDB2.conv5.0.weight',
|
||||
'model.1.sub.21.RDB2.conv5.0.bias', 'model.1.sub.21.RDB3.conv1.0.weight',
|
||||
'model.1.sub.21.RDB3.conv1.0.bias', 'model.1.sub.21.RDB3.conv2.0.weight',
|
||||
'model.1.sub.21.RDB3.conv2.0.bias', 'model.1.sub.21.RDB3.conv3.0.weight',
|
||||
'model.1.sub.21.RDB3.conv3.0.bias', 'model.1.sub.21.RDB3.conv4.0.weight',
|
||||
'model.1.sub.21.RDB3.conv4.0.bias', 'model.1.sub.21.RDB3.conv5.0.weight',
|
||||
'model.1.sub.21.RDB3.conv5.0.bias', 'model.1.sub.22.RDB1.conv1.0.weight',
|
||||
'model.1.sub.22.RDB1.conv1.0.bias', 'model.1.sub.22.RDB1.conv2.0.weight',
|
||||
'model.1.sub.22.RDB1.conv2.0.bias', 'model.1.sub.22.RDB1.conv3.0.weight',
|
||||
'model.1.sub.22.RDB1.conv3.0.bias', 'model.1.sub.22.RDB1.conv4.0.weight',
|
||||
'model.1.sub.22.RDB1.conv4.0.bias', 'model.1.sub.22.RDB1.conv5.0.weight',
|
||||
'model.1.sub.22.RDB1.conv5.0.bias', 'model.1.sub.22.RDB2.conv1.0.weight',
|
||||
'model.1.sub.22.RDB2.conv1.0.bias', 'model.1.sub.22.RDB2.conv2.0.weight',
|
||||
'model.1.sub.22.RDB2.conv2.0.bias', 'model.1.sub.22.RDB2.conv3.0.weight',
|
||||
'model.1.sub.22.RDB2.conv3.0.bias', 'model.1.sub.22.RDB2.conv4.0.weight',
|
||||
'model.1.sub.22.RDB2.conv4.0.bias', 'model.1.sub.22.RDB2.conv5.0.weight',
|
||||
'model.1.sub.22.RDB2.conv5.0.bias', 'model.1.sub.22.RDB3.conv1.0.weight',
|
||||
'model.1.sub.22.RDB3.conv1.0.bias', 'model.1.sub.22.RDB3.conv2.0.weight',
|
||||
'model.1.sub.22.RDB3.conv2.0.bias', 'model.1.sub.22.RDB3.conv3.0.weight',
|
||||
'model.1.sub.22.RDB3.conv3.0.bias', 'model.1.sub.22.RDB3.conv4.0.weight',
|
||||
'model.1.sub.22.RDB3.conv4.0.bias', 'model.1.sub.22.RDB3.conv5.0.weight',
|
||||
'model.1.sub.22.RDB3.conv5.0.bias', 'model.1.sub.23.weight', 'model.1.sub.23.bias',
|
||||
'model.3.weight', 'model.3.bias', 'model.6.weight', 'model.6.bias', 'model.8.weight',
|
||||
'model.8.bias', 'model.10.weight', 'model.10.bias']
|
||||
|
||||
|
||||
# Borrowed from https://github.com/rlaphoenix/VSGAN/blob/master/vsgan/archs/ESRGAN.py
|
||||
# Which enhanced stuff that was already here
|
||||
class RRDBNet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
state_dict,
|
||||
norm=None,
|
||||
act: str = "leakyrelu",
|
||||
upsampler: str = "upconv",
|
||||
mode: B.ConvMode = "CNA",
|
||||
) -> None:
|
||||
"""
|
||||
ESRGAN - Enhanced Super-Resolution Generative Adversarial Networks.
|
||||
By Xintao Wang, Ke Yu, Shixiang Wu, Jinjin Gu, Yihao Liu, Chao Dong, Yu Qiao,
|
||||
and Chen Change Loy.
|
||||
This is old-arch Residual in Residual Dense Block Network and is not
|
||||
the newest revision that's available at github.com/xinntao/ESRGAN.
|
||||
This is on purpose, the newest Network has severely limited the
|
||||
potential use of the Network with no benefits.
|
||||
This network supports model files from both new and old-arch.
|
||||
Args:
|
||||
norm: Normalization layer
|
||||
act: Activation layer
|
||||
upsampler: Upsample layer. upconv, pixel_shuffle
|
||||
mode: Convolution mode
|
||||
"""
|
||||
super(RRDBNet, self).__init__()
|
||||
self.model_arch = "ESRGAN"
|
||||
self.sub_type = "SR"
|
||||
|
||||
self.state = state_dict
|
||||
self.norm = norm
|
||||
self.act = act
|
||||
self.upsampler = upsampler
|
||||
self.mode = mode
|
||||
|
||||
self.state_map = {
|
||||
# currently supports old, new, and newer RRDBNet arch models
|
||||
# ESRGAN, BSRGAN/RealSR, Real-ESRGAN
|
||||
"model.0.weight": ("conv_first.weight",),
|
||||
"model.0.bias": ("conv_first.bias",),
|
||||
"model.1.sub./NB/.weight": ("trunk_conv.weight", "conv_body.weight"),
|
||||
"model.1.sub./NB/.bias": ("trunk_conv.bias", "conv_body.bias"),
|
||||
r"model.1.sub.\1.RDB\2.conv\3.0.\4": (
|
||||
r"RRDB_trunk\.(\d+)\.RDB(\d)\.conv(\d+)\.(weight|bias)",
|
||||
r"body\.(\d+)\.rdb(\d)\.conv(\d+)\.(weight|bias)",
|
||||
),
|
||||
}
|
||||
if "params_ema" in self.state:
|
||||
self.state = self.state["params_ema"]
|
||||
# self.model_arch = "RealESRGAN"
|
||||
self.num_blocks = self.get_num_blocks()
|
||||
self.plus = any("conv1x1" in k for k in self.state.keys())
|
||||
if self.plus:
|
||||
self.model_arch = "ESRGAN+"
|
||||
|
||||
self.state = self.new_to_old_arch(self.state)
|
||||
|
||||
self.key_arr = list(self.state.keys())
|
||||
|
||||
self.in_nc: int = self.state[self.key_arr[0]].shape[1]
|
||||
self.out_nc: int = self.state[self.key_arr[-1]].shape[0]
|
||||
|
||||
self.scale: int = self.get_scale()
|
||||
self.num_filters: int = self.state[self.key_arr[0]].shape[0]
|
||||
|
||||
c2x2 = False
|
||||
if self.state["model.0.weight"].shape[-2] == 2:
|
||||
c2x2 = True
|
||||
self.scale = round(math.sqrt(self.scale / 4))
|
||||
self.model_arch = "ESRGAN-2c2"
|
||||
|
||||
self.supports_fp16 = True
|
||||
self.supports_bfp16 = True
|
||||
self.min_size_restriction = None
|
||||
|
||||
# Detect if pixelunshuffle was used (Real-ESRGAN)
|
||||
if self.in_nc in (self.out_nc * 4, self.out_nc * 16) and self.out_nc in (
|
||||
self.in_nc / 4,
|
||||
self.in_nc / 16,
|
||||
):
|
||||
self.shuffle_factor = int(math.sqrt(self.in_nc / self.out_nc))
|
||||
else:
|
||||
self.shuffle_factor = None
|
||||
|
||||
upsample_block = {
|
||||
"upconv": B.upconv_block,
|
||||
"pixel_shuffle": B.pixelshuffle_block,
|
||||
}.get(self.upsampler)
|
||||
if upsample_block is None:
|
||||
raise NotImplementedError(f"Upsample mode [{self.upsampler}] is not found")
|
||||
|
||||
if self.scale == 3:
|
||||
upsample_blocks = upsample_block(
|
||||
in_nc=self.num_filters,
|
||||
out_nc=self.num_filters,
|
||||
upscale_factor=3,
|
||||
act_type=self.act,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
else:
|
||||
upsample_blocks = [
|
||||
upsample_block(
|
||||
in_nc=self.num_filters,
|
||||
out_nc=self.num_filters,
|
||||
act_type=self.act,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
for _ in range(int(math.log(self.scale, 2)))
|
||||
]
|
||||
|
||||
self.model = B.sequential(
|
||||
# fea conv
|
||||
B.conv_block(
|
||||
in_nc=self.in_nc,
|
||||
out_nc=self.num_filters,
|
||||
kernel_size=3,
|
||||
norm_type=None,
|
||||
act_type=None,
|
||||
c2x2=c2x2,
|
||||
),
|
||||
B.ShortcutBlock(
|
||||
B.sequential(
|
||||
# rrdb blocks
|
||||
*[
|
||||
B.RRDB(
|
||||
nf=self.num_filters,
|
||||
kernel_size=3,
|
||||
gc=32,
|
||||
stride=1,
|
||||
bias=True,
|
||||
pad_type="zero",
|
||||
norm_type=self.norm,
|
||||
act_type=self.act,
|
||||
mode="CNA",
|
||||
plus=self.plus,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
for _ in range(self.num_blocks)
|
||||
],
|
||||
# lr conv
|
||||
B.conv_block(
|
||||
in_nc=self.num_filters,
|
||||
out_nc=self.num_filters,
|
||||
kernel_size=3,
|
||||
norm_type=self.norm,
|
||||
act_type=None,
|
||||
mode=self.mode,
|
||||
c2x2=c2x2,
|
||||
),
|
||||
)
|
||||
),
|
||||
*upsample_blocks,
|
||||
# hr_conv0
|
||||
B.conv_block(
|
||||
in_nc=self.num_filters,
|
||||
out_nc=self.num_filters,
|
||||
kernel_size=3,
|
||||
norm_type=None,
|
||||
act_type=self.act,
|
||||
c2x2=c2x2,
|
||||
),
|
||||
# hr_conv1
|
||||
B.conv_block(
|
||||
in_nc=self.num_filters,
|
||||
out_nc=self.out_nc,
|
||||
kernel_size=3,
|
||||
norm_type=None,
|
||||
act_type=None,
|
||||
c2x2=c2x2,
|
||||
),
|
||||
)
|
||||
|
||||
# Adjust these properties for calculations outside of the model
|
||||
if self.shuffle_factor:
|
||||
self.in_nc //= self.shuffle_factor ** 2
|
||||
self.scale //= self.shuffle_factor
|
||||
|
||||
self.load_state_dict(self.state, strict=False)
|
||||
|
||||
def new_to_old_arch(self, state):
|
||||
"""Convert a new-arch model state dictionary to an old-arch dictionary."""
|
||||
if "params_ema" in state:
|
||||
state = state["params_ema"]
|
||||
|
||||
if "conv_first.weight" not in state:
|
||||
# model is already old arch, this is a loose check, but should be sufficient
|
||||
return state
|
||||
|
||||
# add nb to state keys
|
||||
for kind in ("weight", "bias"):
|
||||
self.state_map[f"model.1.sub.{self.num_blocks}.{kind}"] = self.state_map[
|
||||
f"model.1.sub./NB/.{kind}"
|
||||
]
|
||||
del self.state_map[f"model.1.sub./NB/.{kind}"]
|
||||
|
||||
old_state = OrderedDict()
|
||||
for old_key, new_keys in self.state_map.items():
|
||||
for new_key in new_keys:
|
||||
if r"\1" in old_key:
|
||||
for k, v in state.items():
|
||||
sub = re.sub(new_key, old_key, k)
|
||||
if sub != k:
|
||||
old_state[sub] = v
|
||||
else:
|
||||
if new_key in state:
|
||||
old_state[old_key] = state[new_key]
|
||||
|
||||
# upconv layers
|
||||
max_upconv = 0
|
||||
for key in state.keys():
|
||||
match = re.match(r"(upconv|conv_up)(\d)\.(weight|bias)", key)
|
||||
if match is not None:
|
||||
_, key_num, key_type = match.groups()
|
||||
old_state[f"model.{int(key_num) * 3}.{key_type}"] = state[key]
|
||||
max_upconv = max(max_upconv, int(key_num) * 3)
|
||||
|
||||
# final layers
|
||||
for key in state.keys():
|
||||
if key in ("HRconv.weight", "conv_hr.weight"):
|
||||
old_state[f"model.{max_upconv + 2}.weight"] = state[key]
|
||||
elif key in ("HRconv.bias", "conv_hr.bias"):
|
||||
old_state[f"model.{max_upconv + 2}.bias"] = state[key]
|
||||
elif key in ("conv_last.weight",):
|
||||
old_state[f"model.{max_upconv + 4}.weight"] = state[key]
|
||||
elif key in ("conv_last.bias",):
|
||||
old_state[f"model.{max_upconv + 4}.bias"] = state[key]
|
||||
|
||||
# Sort by first numeric value of each layer
|
||||
def compare(item1, item2):
|
||||
parts1 = item1.split(".")
|
||||
parts2 = item2.split(".")
|
||||
int1 = int(parts1[1])
|
||||
int2 = int(parts2[1])
|
||||
return int1 - int2
|
||||
|
||||
sorted_keys = sorted(old_state.keys(), key=functools.cmp_to_key(compare))
|
||||
|
||||
# Rebuild the output dict in the right order
|
||||
out_dict = OrderedDict((k, old_state[k]) for k in sorted_keys)
|
||||
|
||||
return out_dict
|
||||
|
||||
def get_scale(self, min_part: int = 6) -> int:
|
||||
n = 0
|
||||
for part in list(self.state):
|
||||
parts = part.split(".")[1:]
|
||||
if len(parts) == 2:
|
||||
part_num = int(parts[0])
|
||||
if part_num > min_part and parts[1] == "weight":
|
||||
n += 1
|
||||
return 2 ** n
|
||||
|
||||
def get_num_blocks(self) -> int:
|
||||
nbs = []
|
||||
state_keys = self.state_map[r"model.1.sub.\1.RDB\2.conv\3.0.\4"] + (
|
||||
r"model\.\d+\.sub\.(\d+)\.RDB(\d+)\.conv(\d+)\.0\.(weight|bias)",
|
||||
)
|
||||
for state_key in state_keys:
|
||||
for k in self.state:
|
||||
m = re.search(state_key, k)
|
||||
if m:
|
||||
nbs.append(int(m.group(1)))
|
||||
if nbs:
|
||||
break
|
||||
return max(*nbs) + 1
|
||||
|
||||
def forward(self, x):
|
||||
if self.shuffle_factor:
|
||||
_, _, h, w = x.size()
|
||||
mod_pad_h = (
|
||||
self.shuffle_factor - h % self.shuffle_factor
|
||||
) % self.shuffle_factor
|
||||
mod_pad_w = (
|
||||
self.shuffle_factor - w % self.shuffle_factor
|
||||
) % self.shuffle_factor
|
||||
x = F.pad(x, (0, mod_pad_w, 0, mod_pad_h), "reflect")
|
||||
x = torch.pixel_unshuffle(x, downscale_factor=self.shuffle_factor)
|
||||
x = self.model(x)
|
||||
return x[:, :, : h * self.scale, : w * self.scale]
|
||||
return self.model(x)
|
||||
549
toolkit/models/block.py
Normal file
549
toolkit/models/block.py
Normal file
@@ -0,0 +1,549 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
try:
|
||||
from typing import Literal
|
||||
except ImportError:
|
||||
from typing_extensions import Literal
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
####################
|
||||
# Basic blocks
|
||||
####################
|
||||
|
||||
|
||||
def act(act_type: str, inplace=True, neg_slope=0.2, n_prelu=1):
|
||||
# helper selecting activation
|
||||
# neg_slope: for leakyrelu and init of prelu
|
||||
# n_prelu: for p_relu num_parameters
|
||||
act_type = act_type.lower()
|
||||
if act_type == "relu":
|
||||
layer = nn.ReLU(inplace)
|
||||
elif act_type == "leakyrelu":
|
||||
layer = nn.LeakyReLU(neg_slope, inplace)
|
||||
elif act_type == "prelu":
|
||||
layer = nn.PReLU(num_parameters=n_prelu, init=neg_slope)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"activation layer [{:s}] is not found".format(act_type)
|
||||
)
|
||||
return layer
|
||||
|
||||
|
||||
def norm(norm_type: str, nc: int):
|
||||
# helper selecting normalization layer
|
||||
norm_type = norm_type.lower()
|
||||
if norm_type == "batch":
|
||||
layer = nn.BatchNorm2d(nc, affine=True)
|
||||
elif norm_type == "instance":
|
||||
layer = nn.InstanceNorm2d(nc, affine=False)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"normalization layer [{:s}] is not found".format(norm_type)
|
||||
)
|
||||
return layer
|
||||
|
||||
|
||||
def pad(pad_type: str, padding):
|
||||
# helper selecting padding layer
|
||||
# if padding is 'zero', do by conv layers
|
||||
pad_type = pad_type.lower()
|
||||
if padding == 0:
|
||||
return None
|
||||
if pad_type == "reflect":
|
||||
layer = nn.ReflectionPad2d(padding)
|
||||
elif pad_type == "replicate":
|
||||
layer = nn.ReplicationPad2d(padding)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"padding layer [{:s}] is not implemented".format(pad_type)
|
||||
)
|
||||
return layer
|
||||
|
||||
|
||||
def get_valid_padding(kernel_size, dilation):
|
||||
kernel_size = kernel_size + (kernel_size - 1) * (dilation - 1)
|
||||
padding = (kernel_size - 1) // 2
|
||||
return padding
|
||||
|
||||
|
||||
class ConcatBlock(nn.Module):
|
||||
# Concat the output of a submodule to its input
|
||||
def __init__(self, submodule):
|
||||
super(ConcatBlock, self).__init__()
|
||||
self.sub = submodule
|
||||
|
||||
def forward(self, x):
|
||||
output = torch.cat((x, self.sub(x)), dim=1)
|
||||
return output
|
||||
|
||||
def __repr__(self):
|
||||
tmpstr = "Identity .. \n|"
|
||||
modstr = self.sub.__repr__().replace("\n", "\n|")
|
||||
tmpstr = tmpstr + modstr
|
||||
return tmpstr
|
||||
|
||||
|
||||
class ShortcutBlock(nn.Module):
|
||||
# Elementwise sum the output of a submodule to its input
|
||||
def __init__(self, submodule):
|
||||
super(ShortcutBlock, self).__init__()
|
||||
self.sub = submodule
|
||||
|
||||
def forward(self, x):
|
||||
output = x + self.sub(x)
|
||||
return output
|
||||
|
||||
def __repr__(self):
|
||||
tmpstr = "Identity + \n|"
|
||||
modstr = self.sub.__repr__().replace("\n", "\n|")
|
||||
tmpstr = tmpstr + modstr
|
||||
return tmpstr
|
||||
|
||||
|
||||
class ShortcutBlockSPSR(nn.Module):
|
||||
# Elementwise sum the output of a submodule to its input
|
||||
def __init__(self, submodule):
|
||||
super(ShortcutBlockSPSR, self).__init__()
|
||||
self.sub = submodule
|
||||
|
||||
def forward(self, x):
|
||||
return x, self.sub
|
||||
|
||||
def __repr__(self):
|
||||
tmpstr = "Identity + \n|"
|
||||
modstr = self.sub.__repr__().replace("\n", "\n|")
|
||||
tmpstr = tmpstr + modstr
|
||||
return tmpstr
|
||||
|
||||
|
||||
def sequential(*args):
|
||||
# Flatten Sequential. It unwraps nn.Sequential.
|
||||
if len(args) == 1:
|
||||
if isinstance(args[0], OrderedDict):
|
||||
raise NotImplementedError("sequential does not support OrderedDict input.")
|
||||
return args[0] # No sequential is needed.
|
||||
modules = []
|
||||
for module in args:
|
||||
if isinstance(module, nn.Sequential):
|
||||
for submodule in module.children():
|
||||
modules.append(submodule)
|
||||
elif isinstance(module, nn.Module):
|
||||
modules.append(module)
|
||||
return nn.Sequential(*modules)
|
||||
|
||||
|
||||
ConvMode = Literal["CNA", "NAC", "CNAC"]
|
||||
|
||||
|
||||
# 2x2x2 Conv Block
|
||||
def conv_block_2c2(
|
||||
in_nc,
|
||||
out_nc,
|
||||
act_type="relu",
|
||||
):
|
||||
return sequential(
|
||||
nn.Conv2d(in_nc, out_nc, kernel_size=2, padding=1),
|
||||
nn.Conv2d(out_nc, out_nc, kernel_size=2, padding=0),
|
||||
act(act_type) if act_type else None,
|
||||
)
|
||||
|
||||
|
||||
def conv_block(
|
||||
in_nc: int,
|
||||
out_nc: int,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True,
|
||||
pad_type="zero",
|
||||
norm_type: str | None = None,
|
||||
act_type: str | None = "relu",
|
||||
mode: ConvMode = "CNA",
|
||||
c2x2=False,
|
||||
):
|
||||
"""
|
||||
Conv layer with padding, normalization, activation
|
||||
mode: CNA --> Conv -> Norm -> Act
|
||||
NAC --> Norm -> Act --> Conv (Identity Mappings in Deep Residual Networks, ECCV16)
|
||||
"""
|
||||
|
||||
if c2x2:
|
||||
return conv_block_2c2(in_nc, out_nc, act_type=act_type)
|
||||
|
||||
assert mode in ("CNA", "NAC", "CNAC"), "Wrong conv mode [{:s}]".format(mode)
|
||||
padding = get_valid_padding(kernel_size, dilation)
|
||||
p = pad(pad_type, padding) if pad_type and pad_type != "zero" else None
|
||||
padding = padding if pad_type == "zero" else 0
|
||||
|
||||
c = nn.Conv2d(
|
||||
in_nc,
|
||||
out_nc,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
bias=bias,
|
||||
groups=groups,
|
||||
)
|
||||
a = act(act_type) if act_type else None
|
||||
if mode in ("CNA", "CNAC"):
|
||||
n = norm(norm_type, out_nc) if norm_type else None
|
||||
return sequential(p, c, n, a)
|
||||
elif mode == "NAC":
|
||||
if norm_type is None and act_type is not None:
|
||||
a = act(act_type, inplace=False)
|
||||
# Important!
|
||||
# input----ReLU(inplace)----Conv--+----output
|
||||
# |________________________|
|
||||
# inplace ReLU will modify the input, therefore wrong output
|
||||
n = norm(norm_type, in_nc) if norm_type else None
|
||||
return sequential(n, a, p, c)
|
||||
else:
|
||||
assert False, f"Invalid conv mode {mode}"
|
||||
|
||||
|
||||
####################
|
||||
# Useful blocks
|
||||
####################
|
||||
|
||||
|
||||
class ResNetBlock(nn.Module):
|
||||
"""
|
||||
ResNet Block, 3-3 style
|
||||
with extra residual scaling used in EDSR
|
||||
(Enhanced Deep Residual Networks for Single Image Super-Resolution, CVPRW 17)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_nc,
|
||||
mid_nc,
|
||||
out_nc,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True,
|
||||
pad_type="zero",
|
||||
norm_type=None,
|
||||
act_type="relu",
|
||||
mode: ConvMode = "CNA",
|
||||
res_scale=1,
|
||||
):
|
||||
super(ResNetBlock, self).__init__()
|
||||
conv0 = conv_block(
|
||||
in_nc,
|
||||
mid_nc,
|
||||
kernel_size,
|
||||
stride,
|
||||
dilation,
|
||||
groups,
|
||||
bias,
|
||||
pad_type,
|
||||
norm_type,
|
||||
act_type,
|
||||
mode,
|
||||
)
|
||||
if mode == "CNA":
|
||||
act_type = None
|
||||
if mode == "CNAC": # Residual path: |-CNAC-|
|
||||
act_type = None
|
||||
norm_type = None
|
||||
conv1 = conv_block(
|
||||
mid_nc,
|
||||
out_nc,
|
||||
kernel_size,
|
||||
stride,
|
||||
dilation,
|
||||
groups,
|
||||
bias,
|
||||
pad_type,
|
||||
norm_type,
|
||||
act_type,
|
||||
mode,
|
||||
)
|
||||
# if in_nc != out_nc:
|
||||
# self.project = conv_block(in_nc, out_nc, 1, stride, dilation, 1, bias, pad_type, \
|
||||
# None, None)
|
||||
# print('Need a projecter in ResNetBlock.')
|
||||
# else:
|
||||
# self.project = lambda x:x
|
||||
self.res = sequential(conv0, conv1)
|
||||
self.res_scale = res_scale
|
||||
|
||||
def forward(self, x):
|
||||
res = self.res(x).mul(self.res_scale)
|
||||
return x + res
|
||||
|
||||
|
||||
class RRDB(nn.Module):
|
||||
"""
|
||||
Residual in Residual Dense Block
|
||||
(ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
nf,
|
||||
kernel_size=3,
|
||||
gc=32,
|
||||
stride=1,
|
||||
bias: bool = True,
|
||||
pad_type="zero",
|
||||
norm_type=None,
|
||||
act_type="leakyrelu",
|
||||
mode: ConvMode = "CNA",
|
||||
_convtype="Conv2D",
|
||||
_spectral_norm=False,
|
||||
plus=False,
|
||||
c2x2=False,
|
||||
):
|
||||
super(RRDB, self).__init__()
|
||||
self.RDB1 = ResidualDenseBlock_5C(
|
||||
nf,
|
||||
kernel_size,
|
||||
gc,
|
||||
stride,
|
||||
bias,
|
||||
pad_type,
|
||||
norm_type,
|
||||
act_type,
|
||||
mode,
|
||||
plus=plus,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
self.RDB2 = ResidualDenseBlock_5C(
|
||||
nf,
|
||||
kernel_size,
|
||||
gc,
|
||||
stride,
|
||||
bias,
|
||||
pad_type,
|
||||
norm_type,
|
||||
act_type,
|
||||
mode,
|
||||
plus=plus,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
self.RDB3 = ResidualDenseBlock_5C(
|
||||
nf,
|
||||
kernel_size,
|
||||
gc,
|
||||
stride,
|
||||
bias,
|
||||
pad_type,
|
||||
norm_type,
|
||||
act_type,
|
||||
mode,
|
||||
plus=plus,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
out = self.RDB1(x)
|
||||
out = self.RDB2(out)
|
||||
out = self.RDB3(out)
|
||||
return out * 0.2 + x
|
||||
|
||||
|
||||
class ResidualDenseBlock_5C(nn.Module):
|
||||
"""
|
||||
Residual Dense Block
|
||||
style: 5 convs
|
||||
The core module of paper: (Residual Dense Network for Image Super-Resolution, CVPR 18)
|
||||
Modified options that can be used:
|
||||
- "Partial Convolution based Padding" arXiv:1811.11718
|
||||
- "Spectral normalization" arXiv:1802.05957
|
||||
- "ICASSP 2020 - ESRGAN+ : Further Improving ESRGAN" N. C.
|
||||
{Rakotonirina} and A. {Rasoanaivo}
|
||||
|
||||
Args:
|
||||
nf (int): Channel number of intermediate features (num_feat).
|
||||
gc (int): Channels for each growth (num_grow_ch: growth channel,
|
||||
i.e. intermediate channels).
|
||||
convtype (str): the type of convolution to use. Default: 'Conv2D'
|
||||
gaussian_noise (bool): enable the ESRGAN+ gaussian noise (no new
|
||||
trainable parameters)
|
||||
plus (bool): enable the additional residual paths from ESRGAN+
|
||||
(adds trainable parameters)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
nf=64,
|
||||
kernel_size=3,
|
||||
gc=32,
|
||||
stride=1,
|
||||
bias: bool = True,
|
||||
pad_type="zero",
|
||||
norm_type=None,
|
||||
act_type="leakyrelu",
|
||||
mode: ConvMode = "CNA",
|
||||
plus=False,
|
||||
c2x2=False,
|
||||
):
|
||||
super(ResidualDenseBlock_5C, self).__init__()
|
||||
|
||||
## +
|
||||
self.conv1x1 = conv1x1(nf, gc) if plus else None
|
||||
## +
|
||||
|
||||
self.conv1 = conv_block(
|
||||
nf,
|
||||
gc,
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=act_type,
|
||||
mode=mode,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
self.conv2 = conv_block(
|
||||
nf + gc,
|
||||
gc,
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=act_type,
|
||||
mode=mode,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
self.conv3 = conv_block(
|
||||
nf + 2 * gc,
|
||||
gc,
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=act_type,
|
||||
mode=mode,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
self.conv4 = conv_block(
|
||||
nf + 3 * gc,
|
||||
gc,
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=act_type,
|
||||
mode=mode,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
if mode == "CNA":
|
||||
last_act = None
|
||||
else:
|
||||
last_act = act_type
|
||||
self.conv5 = conv_block(
|
||||
nf + 4 * gc,
|
||||
nf,
|
||||
3,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=last_act,
|
||||
mode=mode,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.conv1(x)
|
||||
x2 = self.conv2(torch.cat((x, x1), 1))
|
||||
if self.conv1x1:
|
||||
# pylint: disable=not-callable
|
||||
x2 = x2 + self.conv1x1(x) # +
|
||||
x3 = self.conv3(torch.cat((x, x1, x2), 1))
|
||||
x4 = self.conv4(torch.cat((x, x1, x2, x3), 1))
|
||||
if self.conv1x1:
|
||||
x4 = x4 + x2 # +
|
||||
x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1))
|
||||
return x5 * 0.2 + x
|
||||
|
||||
|
||||
def conv1x1(in_planes, out_planes, stride=1):
|
||||
return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
|
||||
|
||||
|
||||
####################
|
||||
# Upsampler
|
||||
####################
|
||||
|
||||
|
||||
def pixelshuffle_block(
|
||||
in_nc: int,
|
||||
out_nc: int,
|
||||
upscale_factor=2,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
bias=True,
|
||||
pad_type="zero",
|
||||
norm_type: str | None = None,
|
||||
act_type="relu",
|
||||
):
|
||||
"""
|
||||
Pixel shuffle layer
|
||||
(Real-Time Single Image and Video Super-Resolution Using an Efficient Sub-Pixel Convolutional
|
||||
Neural Network, CVPR17)
|
||||
"""
|
||||
conv = conv_block(
|
||||
in_nc,
|
||||
out_nc * (upscale_factor ** 2),
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=None,
|
||||
act_type=None,
|
||||
)
|
||||
pixel_shuffle = nn.PixelShuffle(upscale_factor)
|
||||
|
||||
n = norm(norm_type, out_nc) if norm_type else None
|
||||
a = act(act_type) if act_type else None
|
||||
return sequential(conv, pixel_shuffle, n, a)
|
||||
|
||||
|
||||
def upconv_block(
|
||||
in_nc: int,
|
||||
out_nc: int,
|
||||
upscale_factor=2,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
bias=True,
|
||||
pad_type="zero",
|
||||
norm_type: str | None = None,
|
||||
act_type="relu",
|
||||
mode="nearest",
|
||||
c2x2=False,
|
||||
):
|
||||
# Up conv
|
||||
# described in https://distill.pub/2016/deconv-checkerboard/
|
||||
# convert to float 16 if is bfloat16
|
||||
upsample = nn.Upsample(scale_factor=upscale_factor, mode=mode)
|
||||
conv = conv_block(
|
||||
in_nc,
|
||||
out_nc,
|
||||
kernel_size,
|
||||
stride,
|
||||
bias=bias,
|
||||
pad_type=pad_type,
|
||||
norm_type=norm_type,
|
||||
act_type=act_type,
|
||||
c2x2=c2x2,
|
||||
)
|
||||
return sequential(upsample, conv)
|
||||
@@ -5,6 +5,12 @@ CONFIG_ROOT = os.path.join(TOOLKIT_ROOT, 'config')
|
||||
SD_SCRIPTS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories", "sd-scripts")
|
||||
REPOS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories")
|
||||
|
||||
# check if ENV variable is set
|
||||
if 'MODELS_PATH' in os.environ:
|
||||
MODELS_PATH = os.environ['MODELS_PATH']
|
||||
else:
|
||||
MODELS_PATH = os.path.join(TOOLKIT_ROOT, "models")
|
||||
|
||||
|
||||
def get_path(path):
|
||||
# we allow absolute paths, but if it is not absolute, we assume it is relative to the toolkit root
|
||||
|
||||
@@ -2,9 +2,76 @@ from typing import Union, List, Optional, Dict, Any, Tuple, Callable
|
||||
|
||||
import torch
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionPipeline
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from diffusers.pipelines.stable_diffusion_xl.watermark import StableDiffusionXLWatermarker
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from diffusers import AutoencoderKL, UNet2DConditionModel
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextModelWithProjection
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
|
||||
|
||||
class FakeWatermarker(StableDiffusionXLWatermarker):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def apply_watermark(self, image):
|
||||
return image
|
||||
|
||||
|
||||
class HackedStableDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
def __init__(
|
||||
self,
|
||||
vae: 'AutoencoderKL',
|
||||
text_encoder: 'CLIPTextModel',
|
||||
text_encoder_2: 'CLIPTextModelWithProjection',
|
||||
tokenizer: 'CLIPTokenizer',
|
||||
tokenizer_2: 'CLIPTokenizer',
|
||||
unet: 'UNet2DConditionModel',
|
||||
scheduler: 'KarrasDiffusionSchedulers',
|
||||
force_zeros_for_empty_prompt: bool = True,
|
||||
add_watermarker: bool = False,
|
||||
):
|
||||
# call parents parent super skipping parent
|
||||
super(StableDiffusionXLPipeline, self).__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer=tokenizer,
|
||||
tokenizer_2=tokenizer_2,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
self.register_to_config(force_zeros_for_empty_prompt=force_zeros_for_empty_prompt)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
self.default_sample_size = 1024
|
||||
|
||||
self.watermark = FakeWatermarker()
|
||||
|
||||
def _get_add_time_ids(self, original_size, crops_coords_top_left, target_size, dtype):
|
||||
add_time_ids = list(original_size + crops_coords_top_left + target_size)
|
||||
|
||||
# passed_add_embed_dim = (
|
||||
# self.unet.config.addition_time_embed_dim * len(add_time_ids) + self.text_encoder_2.config.projection_dim
|
||||
# )
|
||||
# expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features
|
||||
#
|
||||
# if expected_add_embed_dim != passed_add_embed_dim:
|
||||
# raise ValueError(
|
||||
# f"Model expects an added time embedding vector of length {expected_add_embed_dim}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
# )
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
return add_time_ids
|
||||
|
||||
|
||||
class CustomStableDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
# def __init__(self, *args, **kwargs):
|
||||
|
||||
388
toolkit/prompt_utils.py
Normal file
388
toolkit/prompt_utils.py
Normal file
@@ -0,0 +1,388 @@
|
||||
import os
|
||||
from typing import Optional, TYPE_CHECKING, List
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
class ACTION_TYPES_SLIDER:
|
||||
ERASE_NEGATIVE = 0
|
||||
ENHANCE_NEGATIVE = 1
|
||||
|
||||
|
||||
class EncodedPromptPair:
|
||||
def __init__(
|
||||
self,
|
||||
target_class,
|
||||
target_class_with_neutral,
|
||||
positive_target,
|
||||
positive_target_with_neutral,
|
||||
negative_target,
|
||||
negative_target_with_neutral,
|
||||
neutral,
|
||||
empty_prompt,
|
||||
both_targets,
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
action_list=None,
|
||||
multiplier=1.0,
|
||||
multiplier_list=None,
|
||||
weight=1.0
|
||||
):
|
||||
self.target_class: PromptEmbeds = target_class
|
||||
self.target_class_with_neutral: PromptEmbeds = target_class_with_neutral
|
||||
self.positive_target: PromptEmbeds = positive_target
|
||||
self.positive_target_with_neutral: PromptEmbeds = positive_target_with_neutral
|
||||
self.negative_target: PromptEmbeds = negative_target
|
||||
self.negative_target_with_neutral: PromptEmbeds = negative_target_with_neutral
|
||||
self.neutral: PromptEmbeds = neutral
|
||||
self.empty_prompt: PromptEmbeds = empty_prompt
|
||||
self.both_targets: PromptEmbeds = both_targets
|
||||
self.multiplier: float = multiplier
|
||||
if multiplier_list is not None:
|
||||
self.multiplier_list: list[float] = multiplier_list
|
||||
else:
|
||||
self.multiplier_list: list[float] = [multiplier]
|
||||
self.action: int = action
|
||||
if action_list is not None:
|
||||
self.action_list: list[int] = action_list
|
||||
else:
|
||||
self.action_list: list[int] = [action]
|
||||
self.weight: float = weight
|
||||
|
||||
# simulate torch to for tensors
|
||||
def to(self, *args, **kwargs):
|
||||
self.target_class = self.target_class.to(*args, **kwargs)
|
||||
self.target_class_with_neutral = self.target_class_with_neutral.to(*args, **kwargs)
|
||||
self.positive_target = self.positive_target.to(*args, **kwargs)
|
||||
self.positive_target_with_neutral = self.positive_target_with_neutral.to(*args, **kwargs)
|
||||
self.negative_target = self.negative_target.to(*args, **kwargs)
|
||||
self.negative_target_with_neutral = self.negative_target_with_neutral.to(*args, **kwargs)
|
||||
self.neutral = self.neutral.to(*args, **kwargs)
|
||||
self.empty_prompt = self.empty_prompt.to(*args, **kwargs)
|
||||
self.both_targets = self.both_targets.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
|
||||
def concat_prompt_embeds(prompt_embeds: list[PromptEmbeds]):
|
||||
text_embeds = torch.cat([p.text_embeds for p in prompt_embeds], dim=0)
|
||||
pooled_embeds = None
|
||||
if prompt_embeds[0].pooled_embeds is not None:
|
||||
pooled_embeds = torch.cat([p.pooled_embeds for p in prompt_embeds], dim=0)
|
||||
return PromptEmbeds([text_embeds, pooled_embeds])
|
||||
|
||||
|
||||
def concat_prompt_pairs(prompt_pairs: list[EncodedPromptPair]):
|
||||
weight = prompt_pairs[0].weight
|
||||
target_class = concat_prompt_embeds([p.target_class for p in prompt_pairs])
|
||||
target_class_with_neutral = concat_prompt_embeds([p.target_class_with_neutral for p in prompt_pairs])
|
||||
positive_target = concat_prompt_embeds([p.positive_target for p in prompt_pairs])
|
||||
positive_target_with_neutral = concat_prompt_embeds([p.positive_target_with_neutral for p in prompt_pairs])
|
||||
negative_target = concat_prompt_embeds([p.negative_target for p in prompt_pairs])
|
||||
negative_target_with_neutral = concat_prompt_embeds([p.negative_target_with_neutral for p in prompt_pairs])
|
||||
neutral = concat_prompt_embeds([p.neutral for p in prompt_pairs])
|
||||
empty_prompt = concat_prompt_embeds([p.empty_prompt for p in prompt_pairs])
|
||||
both_targets = concat_prompt_embeds([p.both_targets for p in prompt_pairs])
|
||||
# combine all the lists
|
||||
action_list = []
|
||||
multiplier_list = []
|
||||
weight_list = []
|
||||
for p in prompt_pairs:
|
||||
action_list += p.action_list
|
||||
multiplier_list += p.multiplier_list
|
||||
return EncodedPromptPair(
|
||||
target_class=target_class,
|
||||
target_class_with_neutral=target_class_with_neutral,
|
||||
positive_target=positive_target,
|
||||
positive_target_with_neutral=positive_target_with_neutral,
|
||||
negative_target=negative_target,
|
||||
negative_target_with_neutral=negative_target_with_neutral,
|
||||
neutral=neutral,
|
||||
empty_prompt=empty_prompt,
|
||||
both_targets=both_targets,
|
||||
action_list=action_list,
|
||||
multiplier_list=multiplier_list,
|
||||
weight=weight
|
||||
)
|
||||
|
||||
|
||||
def split_prompt_embeds(concatenated: PromptEmbeds, num_parts=None) -> List[PromptEmbeds]:
|
||||
if num_parts is None:
|
||||
# use batch size
|
||||
num_parts = concatenated.text_embeds.shape[0]
|
||||
text_embeds_splits = torch.chunk(concatenated.text_embeds, num_parts, dim=0)
|
||||
|
||||
if concatenated.pooled_embeds is not None:
|
||||
pooled_embeds_splits = torch.chunk(concatenated.pooled_embeds, num_parts, dim=0)
|
||||
else:
|
||||
pooled_embeds_splits = [None] * num_parts
|
||||
|
||||
prompt_embeds_list = [
|
||||
PromptEmbeds([text, pooled])
|
||||
for text, pooled in zip(text_embeds_splits, pooled_embeds_splits)
|
||||
]
|
||||
|
||||
return prompt_embeds_list
|
||||
|
||||
|
||||
def split_prompt_pairs(concatenated: EncodedPromptPair, num_embeds=None) -> List[EncodedPromptPair]:
|
||||
target_class_splits = split_prompt_embeds(concatenated.target_class, num_embeds)
|
||||
target_class_with_neutral_splits = split_prompt_embeds(concatenated.target_class_with_neutral, num_embeds)
|
||||
positive_target_splits = split_prompt_embeds(concatenated.positive_target, num_embeds)
|
||||
positive_target_with_neutral_splits = split_prompt_embeds(concatenated.positive_target_with_neutral, num_embeds)
|
||||
negative_target_splits = split_prompt_embeds(concatenated.negative_target, num_embeds)
|
||||
negative_target_with_neutral_splits = split_prompt_embeds(concatenated.negative_target_with_neutral, num_embeds)
|
||||
neutral_splits = split_prompt_embeds(concatenated.neutral, num_embeds)
|
||||
empty_prompt_splits = split_prompt_embeds(concatenated.empty_prompt, num_embeds)
|
||||
both_targets_splits = split_prompt_embeds(concatenated.both_targets, num_embeds)
|
||||
|
||||
prompt_pairs = []
|
||||
for i in range(len(target_class_splits)):
|
||||
action_list_split = concatenated.action_list[i::len(target_class_splits)]
|
||||
multiplier_list_split = concatenated.multiplier_list[i::len(target_class_splits)]
|
||||
|
||||
prompt_pair = EncodedPromptPair(
|
||||
target_class=target_class_splits[i],
|
||||
target_class_with_neutral=target_class_with_neutral_splits[i],
|
||||
positive_target=positive_target_splits[i],
|
||||
positive_target_with_neutral=positive_target_with_neutral_splits[i],
|
||||
negative_target=negative_target_splits[i],
|
||||
negative_target_with_neutral=negative_target_with_neutral_splits[i],
|
||||
neutral=neutral_splits[i],
|
||||
empty_prompt=empty_prompt_splits[i],
|
||||
both_targets=both_targets_splits[i],
|
||||
action_list=action_list_split,
|
||||
multiplier_list=multiplier_list_split,
|
||||
weight=concatenated.weight
|
||||
)
|
||||
prompt_pairs.append(prompt_pair)
|
||||
|
||||
return prompt_pairs
|
||||
|
||||
|
||||
class PromptEmbedsCache:
|
||||
prompts: dict[str, PromptEmbeds] = {}
|
||||
|
||||
def __setitem__(self, __name: str, __value: PromptEmbeds) -> None:
|
||||
self.prompts[__name] = __value
|
||||
|
||||
def __getitem__(self, __name: str) -> Optional[PromptEmbeds]:
|
||||
if __name in self.prompts:
|
||||
return self.prompts[__name]
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
class EncodedAnchor:
|
||||
def __init__(
|
||||
self,
|
||||
prompt,
|
||||
neg_prompt,
|
||||
multiplier=1.0,
|
||||
multiplier_list=None
|
||||
):
|
||||
self.prompt = prompt
|
||||
self.neg_prompt = neg_prompt
|
||||
self.multiplier = multiplier
|
||||
|
||||
if multiplier_list is not None:
|
||||
self.multiplier_list: list[float] = multiplier_list
|
||||
else:
|
||||
self.multiplier_list: list[float] = [multiplier]
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.prompt = self.prompt.to(*args, **kwargs)
|
||||
self.neg_prompt = self.neg_prompt.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
|
||||
def concat_anchors(anchors: list[EncodedAnchor]):
|
||||
prompt = concat_prompt_embeds([a.prompt for a in anchors])
|
||||
neg_prompt = concat_prompt_embeds([a.neg_prompt for a in anchors])
|
||||
return EncodedAnchor(
|
||||
prompt=prompt,
|
||||
neg_prompt=neg_prompt,
|
||||
multiplier_list=[a.multiplier for a in anchors]
|
||||
)
|
||||
|
||||
|
||||
def split_anchors(concatenated: EncodedAnchor, num_anchors: int = 4) -> List[EncodedAnchor]:
|
||||
prompt_splits = split_prompt_embeds(concatenated.prompt, num_anchors)
|
||||
neg_prompt_splits = split_prompt_embeds(concatenated.neg_prompt, num_anchors)
|
||||
multiplier_list_splits = torch.chunk(torch.tensor(concatenated.multiplier_list), num_anchors)
|
||||
|
||||
anchors = []
|
||||
for prompt, neg_prompt, multiplier in zip(prompt_splits, neg_prompt_splits, multiplier_list_splits):
|
||||
anchor = EncodedAnchor(
|
||||
prompt=prompt,
|
||||
neg_prompt=neg_prompt,
|
||||
multiplier=multiplier.tolist()
|
||||
)
|
||||
anchors.append(anchor)
|
||||
|
||||
return anchors
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_prompts_to_cache(
|
||||
prompt_list: list[str],
|
||||
sd: "StableDiffusion",
|
||||
cache: Optional[PromptEmbedsCache] = None,
|
||||
prompt_tensor_file: Optional[str] = None,
|
||||
) -> PromptEmbedsCache:
|
||||
# TODO: add support for larger prompts
|
||||
if cache is None:
|
||||
cache = PromptEmbedsCache()
|
||||
|
||||
if prompt_tensor_file is not None:
|
||||
# check to see if it exists
|
||||
if os.path.exists(prompt_tensor_file):
|
||||
# load it.
|
||||
print(f"Loading prompt tensors from {prompt_tensor_file}")
|
||||
prompt_tensors = load_file(prompt_tensor_file, device='cpu')
|
||||
# add them to the cache
|
||||
for prompt_txt, prompt_tensor in tqdm(prompt_tensors.items(), desc="Loading prompts", leave=False):
|
||||
if prompt_txt.startswith("te:"):
|
||||
prompt = prompt_txt[3:]
|
||||
# text_embeds
|
||||
text_embeds = prompt_tensor
|
||||
pooled_embeds = None
|
||||
# find pool embeds
|
||||
if f"pe:{prompt}" in prompt_tensors:
|
||||
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
|
||||
|
||||
# make it
|
||||
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
|
||||
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
|
||||
|
||||
if len(cache.prompts) == 0:
|
||||
print("Prompt tensors not found. Encoding prompts..")
|
||||
empty_prompt = ""
|
||||
# encode empty_prompt
|
||||
cache[empty_prompt] = sd.encode_prompt(empty_prompt)
|
||||
|
||||
for p in tqdm(prompt_list, desc="Encoding prompts", leave=False):
|
||||
# build the cache
|
||||
if cache[p] is None:
|
||||
cache[p] = sd.encode_prompt(p).to(device="cpu", dtype=torch.float16)
|
||||
|
||||
# should we shard? It can get large
|
||||
if prompt_tensor_file:
|
||||
print(f"Saving prompt tensors to {prompt_tensor_file}")
|
||||
state_dict = {}
|
||||
for prompt_txt, prompt_embeds in cache.prompts.items():
|
||||
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to(
|
||||
"cpu", dtype=get_torch_dtype('fp16')
|
||||
)
|
||||
if prompt_embeds.pooled_embeds is not None:
|
||||
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to(
|
||||
"cpu",
|
||||
dtype=get_torch_dtype('fp16')
|
||||
)
|
||||
save_file(state_dict, prompt_tensor_file)
|
||||
|
||||
return cache
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.config_modules import SliderTargetConfig
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def build_prompt_pair_batch_from_cache(
|
||||
cache: PromptEmbedsCache,
|
||||
target: 'SliderTargetConfig',
|
||||
neutral: Optional[str] = '',
|
||||
) -> list[EncodedPromptPair]:
|
||||
erase_negative = len(target.positive.strip()) == 0
|
||||
enhance_positive = len(target.negative.strip()) == 0
|
||||
|
||||
both = not erase_negative and not enhance_positive
|
||||
|
||||
prompt_pair_batch = []
|
||||
|
||||
if both or erase_negative:
|
||||
# print("Encoding erase negative")
|
||||
prompt_pair_batch += [
|
||||
# erase standard
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.positive}"],
|
||||
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
negative_target=cache[f"{target.negative}"],
|
||||
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
multiplier=target.multiplier,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
if both or enhance_positive:
|
||||
# print("Encoding enhance positive")
|
||||
prompt_pair_batch += [
|
||||
# enhance standard, swap pos neg
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.negative}"],
|
||||
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
negative_target=cache[f"{target.positive}"],
|
||||
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||
multiplier=target.multiplier,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
if both or enhance_positive:
|
||||
# print("Encoding erase positive (inverse)")
|
||||
prompt_pair_batch += [
|
||||
# erase inverted
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.negative}"],
|
||||
positive_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
negative_target=cache[f"{target.positive}"],
|
||||
negative_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
empty_prompt=cache[""],
|
||||
multiplier=target.multiplier * -1.0,
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
if both or erase_negative:
|
||||
# print("Encoding enhance negative (inverse)")
|
||||
prompt_pair_batch += [
|
||||
# enhance inverted
|
||||
EncodedPromptPair(
|
||||
target_class=cache[target.target_class],
|
||||
target_class_with_neutral=cache[f"{target.target_class} {neutral}"],
|
||||
positive_target=cache[f"{target.positive}"],
|
||||
positive_target_with_neutral=cache[f"{target.positive} {neutral}"],
|
||||
negative_target=cache[f"{target.negative}"],
|
||||
negative_target_with_neutral=cache[f"{target.negative} {neutral}"],
|
||||
both_targets=cache[f"{target.positive} {target.negative}"],
|
||||
neutral=cache[neutral],
|
||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||
empty_prompt=cache[""],
|
||||
multiplier=target.multiplier * -1.0,
|
||||
weight=target.weight
|
||||
),
|
||||
]
|
||||
|
||||
return prompt_pair_batch
|
||||
33
toolkit/scheduler.py
Normal file
33
toolkit/scheduler.py
Normal file
@@ -0,0 +1,33 @@
|
||||
import torch
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def get_lr_scheduler(
|
||||
name: Optional[str],
|
||||
optimizer: torch.optim.Optimizer,
|
||||
max_iterations: Optional[int],
|
||||
lr_min: Optional[float],
|
||||
**kwargs,
|
||||
):
|
||||
if name == "cosine":
|
||||
return torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||
optimizer, T_max=max_iterations, eta_min=lr_min, **kwargs
|
||||
)
|
||||
elif name == "cosine_with_restarts":
|
||||
return torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(
|
||||
optimizer, T_0=max_iterations, T_mult=2, eta_min=lr_min, **kwargs
|
||||
)
|
||||
elif name == "step":
|
||||
return torch.optim.lr_scheduler.StepLR(
|
||||
optimizer, step_size=max_iterations // 100, gamma=0.999, **kwargs
|
||||
)
|
||||
elif name == "constant":
|
||||
return torch.optim.lr_scheduler.ConstantLR(optimizer, factor=1, **kwargs)
|
||||
elif name == "linear":
|
||||
return torch.optim.lr_scheduler.LinearLR(
|
||||
optimizer, start_factor=0.5, end_factor=0.5, total_iters=max_iterations, **kwargs
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Scheduler must be cosine, cosine_with_restarts, step, linear or constant"
|
||||
)
|
||||
@@ -1,26 +1,69 @@
|
||||
import gc
|
||||
import typing
|
||||
from typing import Union, OrderedDict
|
||||
from typing import Union, OrderedDict, List, Tuple
|
||||
import sys
|
||||
import os
|
||||
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||
from safetensors.torch import save_file
|
||||
from torch.amp import autocast
|
||||
from tqdm import tqdm
|
||||
from torchvision.transforms import Resize
|
||||
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from library.sdxl_train_util import _load_target_model as load_sdxl_target_model
|
||||
from toolkit import train_tools
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.model_util_sdxl import load_models_from_sdxl_checkpoint
|
||||
from toolkit.paths import REPOS_ROOT, MODELS_PATH
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'sd-scripts'))
|
||||
from leco import train_util
|
||||
import torch
|
||||
from library import model_util
|
||||
from library.sdxl_model_util import convert_text_encoder_2_state_dict_to_sdxl
|
||||
from library.model_util import convert_unet_state_dict_to_sd, convert_text_encoder_state_dict_to_sd_v2
|
||||
from library import model_util, train_util as kohya_train_util, sdxl_train_util
|
||||
from library.sdxl_model_util import convert_text_encoder_2_state_dict_to_sdxl, DIFFUSERS_SDXL_UNET_CONFIG
|
||||
from diffusers.schedulers import DDPMScheduler
|
||||
from toolkit.pipelines import CustomStableDiffusionXLPipeline, CustomStableDiffusionPipeline, \
|
||||
HackedStableDiffusionXLPipeline
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
import diffusers
|
||||
|
||||
# tell it to shut up
|
||||
diffusers.logging.set_verbosity(diffusers.logging.ERROR)
|
||||
|
||||
|
||||
class BlankNetwork:
|
||||
multiplier = 1.0
|
||||
is_active = True
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __enter__(self):
|
||||
self.is_active = True
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.is_active = False
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも同じ。
|
||||
VAE_SCALE_FACTOR = 8 # 2 ** (len(vae.config.block_out_channels) - 1) = 8
|
||||
|
||||
|
||||
class PromptEmbeds:
|
||||
text_embeds: torch.FloatTensor
|
||||
pooled_embeds: Union[torch.FloatTensor, None]
|
||||
text_embeds: torch.Tensor
|
||||
pooled_embeds: Union[torch.Tensor, None]
|
||||
|
||||
def __init__(self, args) -> None:
|
||||
def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Tensor]) -> None:
|
||||
if isinstance(args, list) or isinstance(args, tuple):
|
||||
# xl
|
||||
self.text_embeds = args[0]
|
||||
@@ -37,33 +80,529 @@ class PromptEmbeds:
|
||||
return self
|
||||
|
||||
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextModelWithProjection
|
||||
|
||||
# if is type checking
|
||||
if typing.TYPE_CHECKING:
|
||||
from diffusers import StableDiffusionPipeline
|
||||
from toolkit.pipelines import CustomStableDiffusionXLPipeline
|
||||
from diffusers import \
|
||||
StableDiffusionPipeline, \
|
||||
AutoencoderKL, \
|
||||
UNet2DConditionModel
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
|
||||
|
||||
class StableDiffusion:
|
||||
pipeline: Union[None, 'StableDiffusionPipeline', 'CustomStableDiffusionXLPipeline']
|
||||
vae: Union[None, 'AutoencoderKL']
|
||||
unet: Union[None, 'UNet2DConditionModel']
|
||||
text_encoder: Union[None, 'CLIPTextModel', List[Union['CLIPTextModel', 'CLIPTextModelWithProjection']]]
|
||||
tokenizer: Union[None, 'CLIPTokenizer', List['CLIPTokenizer']]
|
||||
noise_scheduler: Union[None, 'KarrasDiffusionSchedulers', 'DDPMScheduler']
|
||||
device: str
|
||||
dtype: str
|
||||
torch_dtype: torch.dtype
|
||||
device_torch: torch.device
|
||||
model_config: ModelConfig
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
unet,
|
||||
noise_scheduler,
|
||||
is_xl=False,
|
||||
pipeline=None,
|
||||
device,
|
||||
model_config: ModelConfig,
|
||||
dtype='fp16',
|
||||
custom_pipeline=None,
|
||||
):
|
||||
# text encoder has a list of 2 for xl
|
||||
self.vae = vae
|
||||
self.tokenizer = tokenizer
|
||||
self.text_encoder = text_encoder
|
||||
self.unet = unet
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.is_xl = is_xl
|
||||
self.pipeline = pipeline
|
||||
self.custom_pipeline = custom_pipeline
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(dtype)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.model_config = model_config
|
||||
self.prediction_type = "v_prediction" if self.model_config.is_v_pred else "epsilon"
|
||||
|
||||
# sdxl stuff
|
||||
self.logit_scale = None
|
||||
self.ckpt_info = None
|
||||
self.is_loaded = False
|
||||
|
||||
# to hold network if there is one
|
||||
self.network = None
|
||||
self.is_xl = model_config.is_xl
|
||||
self.is_v2 = model_config.is_v2
|
||||
|
||||
def load_model(self):
|
||||
if self.is_loaded:
|
||||
return
|
||||
dtype = get_torch_dtype(self.dtype)
|
||||
|
||||
# TODO handle other schedulers
|
||||
# sch = KDPM2DiscreteScheduler
|
||||
sch = DDPMScheduler
|
||||
# do our own scheduler
|
||||
prediction_type = "v_prediction" if self.model_config.is_v_pred else "epsilon"
|
||||
scheduler = sch(
|
||||
num_train_timesteps=1000,
|
||||
beta_start=0.00085,
|
||||
beta_end=0.0120,
|
||||
beta_schedule="scaled_linear",
|
||||
clip_sample=False,
|
||||
prediction_type=prediction_type,
|
||||
steps_offset=1
|
||||
)
|
||||
|
||||
model_path = self.model_config.name_or_path
|
||||
if 'civitai.com' in self.model_config.name_or_path:
|
||||
# load is a civit ai model, use the loader.
|
||||
from toolkit.civitai import get_model_path_from_url
|
||||
model_path = get_model_path_from_url(self.model_config.name_or_path)
|
||||
|
||||
if self.model_config.is_xl:
|
||||
# load from kohya
|
||||
# (
|
||||
# load_stable_diffusion_format,
|
||||
# text_encoder1,
|
||||
# text_encoder2,
|
||||
# vae,
|
||||
# unet,
|
||||
# logit_scale,
|
||||
# ckpt_info,
|
||||
# ) = load_sdxl_target_model(
|
||||
# model_path,
|
||||
# self.model_config.vae_path if self.model_config.vae_path is not None else model_path,
|
||||
# 'sdxl',
|
||||
# self.dtype,
|
||||
# self.device,
|
||||
# )
|
||||
|
||||
(
|
||||
text_encoder1,
|
||||
text_encoder2,
|
||||
vae,
|
||||
unet,
|
||||
logit_scale,
|
||||
ckpt_info,
|
||||
) = load_models_from_sdxl_checkpoint(
|
||||
'sdxl',
|
||||
model_path,
|
||||
self.device
|
||||
)
|
||||
|
||||
class Config:
|
||||
def __init__(self):
|
||||
# add all items from DIFFUSERS_SDXL_UNET_CONFIG as attributes
|
||||
for k, v in DIFFUSERS_SDXL_UNET_CONFIG.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
# add diffusers stuff
|
||||
unet.config = Config()
|
||||
|
||||
if self.model_config.vae_path is not None:
|
||||
vae = model_util.load_vae(self.model_config.vae_path, self.dtype)
|
||||
print("additional VAE loaded")
|
||||
|
||||
text_encoder1, text_encoder2, unet = kohya_train_util.transform_models_if_DDP(
|
||||
[text_encoder1, text_encoder2, unet]
|
||||
)
|
||||
class Args:
|
||||
tokenizer_cache_dir = os.path.join(MODELS_PATH, 'clip_cache')
|
||||
|
||||
args = Args()
|
||||
os.makedirs(args.tokenizer_cache_dir, exist_ok=True)
|
||||
|
||||
self.tokenizer = sdxl_train_util.load_tokenizers(args)
|
||||
# tokenizer1 = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
|
||||
# tokenizer2 = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer_2")
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
self.vae = vae
|
||||
self.unet = unet.to(self.device_torch, dtype=dtype)
|
||||
self.text_encoder = [text_encoder1, text_encoder2]
|
||||
self.logit_scale = logit_scale
|
||||
self.ckpt_info = ckpt_info
|
||||
|
||||
|
||||
|
||||
# if self.custom_pipeline is not None:
|
||||
# pipln = self.custom_pipeline
|
||||
# else:
|
||||
# pipln = CustomStableDiffusionXLPipeline
|
||||
#
|
||||
# # see if path exists
|
||||
# if not os.path.exists(model_path):
|
||||
# # try to load with default diffusers
|
||||
# pipe = pipln.from_pretrained(
|
||||
# model_path,
|
||||
# dtype=dtype,
|
||||
# scheduler_type='ddpm',
|
||||
# device=self.device_torch,
|
||||
# ).to(self.device_torch)
|
||||
# else:
|
||||
# pipe = pipln.from_single_file(
|
||||
# model_path,
|
||||
# dtype=dtype,
|
||||
# scheduler_type='ddpm',
|
||||
# device=self.device_torch,
|
||||
# ).to(self.device_torch)
|
||||
#
|
||||
# text_encoders = [pipe.text_encoder, pipe.text_encoder_2]
|
||||
# tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
||||
# for text_encoder in text_encoders:
|
||||
# text_encoder.to(self.device_torch, dtype=dtype)
|
||||
# text_encoder.requires_grad_(False)
|
||||
# text_encoder.eval()
|
||||
# text_encoder = text_encoders
|
||||
else:
|
||||
if self.custom_pipeline is not None:
|
||||
pipln = self.custom_pipeline
|
||||
else:
|
||||
pipln = CustomStableDiffusionPipeline
|
||||
|
||||
# see if path exists
|
||||
if not os.path.exists(model_path):
|
||||
# try to load with default diffusers
|
||||
pipe = pipln.from_pretrained(
|
||||
model_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='dpm',
|
||||
device=self.device_torch,
|
||||
load_safety_checker=False,
|
||||
requires_safety_checker=False,
|
||||
safety_checker=False
|
||||
).to(self.device_torch)
|
||||
else:
|
||||
pipe = pipln.from_single_file(
|
||||
model_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='dpm',
|
||||
device=self.device_torch,
|
||||
load_safety_checker=False,
|
||||
requires_safety_checker=False,
|
||||
safety_checker=False
|
||||
).to(self.device_torch)
|
||||
|
||||
pipe.register_to_config(requires_safety_checker=False)
|
||||
text_encoder = pipe.text_encoder
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
tokenizer = pipe.tokenizer
|
||||
self.tokenizer = tokenizer
|
||||
self.text_encoder = text_encoder
|
||||
self.pipeline = pipe
|
||||
|
||||
# scheduler doesn't get set sometimes, so we set it here
|
||||
pipe.scheduler = scheduler
|
||||
|
||||
self.unet = pipe.unet
|
||||
self.vae = pipe.vae.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.noise_scheduler = scheduler
|
||||
self.vae.eval()
|
||||
self.vae.requires_grad_(False)
|
||||
self.unet.to(self.device_torch, dtype=dtype)
|
||||
self.unet.requires_grad_(False)
|
||||
self.unet.eval()
|
||||
self.is_loaded = True
|
||||
|
||||
def generate_images(self, image_configs: List[GenerateImageConfig]):
|
||||
# sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if self.network is not None:
|
||||
self.network.eval()
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
original_device_dict = {
|
||||
'vae': self.vae.device,
|
||||
'unet': self.unet.device,
|
||||
# 'tokenizer': self.tokenizer.device,
|
||||
}
|
||||
|
||||
# handle sdxl text encoder
|
||||
if isinstance(self.text_encoder, list):
|
||||
for encoder, i in zip(self.text_encoder, range(len(self.text_encoder))):
|
||||
original_device_dict[f'text_encoder_{i}'] = encoder.device
|
||||
encoder.to(self.device_torch)
|
||||
else:
|
||||
original_device_dict['text_encoder'] = self.text_encoder.device
|
||||
self.text_encoder.to(self.device_torch)
|
||||
|
||||
self.vae.to(self.device_torch)
|
||||
self.unet.to(self.device_torch)
|
||||
|
||||
# TODO add clip skip
|
||||
if self.is_xl:
|
||||
pipeline = HackedStableDiffusionXLPipeline(
|
||||
vae=self.vae,
|
||||
unet=self.unet,
|
||||
text_encoder=self.text_encoder[0],
|
||||
text_encoder_2=self.text_encoder[1],
|
||||
tokenizer=self.tokenizer[0],
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
scheduler=self.noise_scheduler,
|
||||
add_watermarker=False,
|
||||
).to(self.device_torch)
|
||||
# force turn that (ruin your images with obvious green and red dots) the #$@@ off!!!
|
||||
pipeline.watermark = None
|
||||
else:
|
||||
pipeline = StableDiffusionPipeline(
|
||||
vae=self.vae,
|
||||
unet=self.unet,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=self.noise_scheduler,
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
).to(self.device_torch)
|
||||
# disable progress bar
|
||||
pipeline.set_progress_bar_config(disable=True)
|
||||
|
||||
start_multiplier = 1.0
|
||||
if self.network is not None:
|
||||
start_multiplier = self.network.multiplier
|
||||
|
||||
pipeline.to(self.device_torch)
|
||||
with network:
|
||||
with torch.no_grad():
|
||||
if self.network is not None:
|
||||
assert self.network.is_active
|
||||
|
||||
for i in tqdm(range(len(image_configs)), desc=f"Generating Images", leave=False):
|
||||
gen_config = image_configs[i]
|
||||
|
||||
if self.network is not None:
|
||||
self.network.multiplier = gen_config.network_multiplier
|
||||
torch.manual_seed(gen_config.seed)
|
||||
torch.cuda.manual_seed(gen_config.seed)
|
||||
|
||||
if self.is_xl:
|
||||
with autocast('cuda'):
|
||||
img = pipeline(
|
||||
prompt=gen_config.prompt,
|
||||
prompt_2=gen_config.prompt_2,
|
||||
negative_prompt=gen_config.negative_prompt,
|
||||
negative_prompt_2=gen_config.negative_prompt_2,
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
guidance_rescale=gen_config.guidance_rescale,
|
||||
).images[0]
|
||||
else:
|
||||
img = pipeline(
|
||||
prompt=gen_config.prompt,
|
||||
negative_prompt=gen_config.negative_prompt,
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
).images[0]
|
||||
|
||||
gen_config.save_image(img)
|
||||
|
||||
# clear pipeline and cache to reduce vram usage
|
||||
del pipeline
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# restore training state
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
|
||||
self.vae.to(original_device_dict['vae'])
|
||||
self.unet.to(original_device_dict['unet'])
|
||||
if isinstance(self.text_encoder, list):
|
||||
for encoder, i in zip(self.text_encoder, range(len(self.text_encoder))):
|
||||
encoder.to(original_device_dict[f'text_encoder_{i}'])
|
||||
else:
|
||||
self.text_encoder.to(original_device_dict['text_encoder'])
|
||||
if self.network is not None:
|
||||
self.network.train()
|
||||
self.network.multiplier = start_multiplier
|
||||
# self.tokenizer.to(original_device_dict['tokenizer'])
|
||||
|
||||
def get_latent_noise(
|
||||
self,
|
||||
height=None,
|
||||
width=None,
|
||||
pixel_height=None,
|
||||
pixel_width=None,
|
||||
batch_size=1,
|
||||
noise_offset=0.0,
|
||||
):
|
||||
if height is None and pixel_height is None:
|
||||
raise ValueError("height or pixel_height must be specified")
|
||||
if width is None and pixel_width is None:
|
||||
raise ValueError("width or pixel_width must be specified")
|
||||
if height is None:
|
||||
height = pixel_height // VAE_SCALE_FACTOR
|
||||
if width is None:
|
||||
width = pixel_width // VAE_SCALE_FACTOR
|
||||
|
||||
noise = torch.randn(
|
||||
(
|
||||
batch_size,
|
||||
UNET_IN_CHANNELS,
|
||||
height,
|
||||
width,
|
||||
),
|
||||
device="cpu",
|
||||
)
|
||||
noise = apply_noise_offset(noise, noise_offset)
|
||||
return noise
|
||||
|
||||
def get_time_ids_from_latents(self, latents: torch.Tensor):
|
||||
bs, ch, h, w = list(latents.shape)
|
||||
|
||||
height = h * VAE_SCALE_FACTOR
|
||||
width = w * VAE_SCALE_FACTOR
|
||||
|
||||
dtype = latents.dtype
|
||||
|
||||
if self.is_xl:
|
||||
prompt_ids = train_util.get_add_time_ids(
|
||||
height,
|
||||
width,
|
||||
dynamic_crops=False, # look into this
|
||||
dtype=dtype,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
return prompt_ids
|
||||
else:
|
||||
return None
|
||||
|
||||
def predict_noise(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
text_embeddings: Union[PromptEmbeds, None] = None,
|
||||
timestep: Union[int, torch.Tensor] = 1,
|
||||
guidance_scale=7.5,
|
||||
guidance_rescale=0, # 0.7 sdxl
|
||||
add_time_ids=None,
|
||||
conditional_embeddings: Union[PromptEmbeds, None] = None,
|
||||
unconditional_embeddings: Union[PromptEmbeds, None] = None,
|
||||
**kwargs,
|
||||
):
|
||||
# get the embeddings
|
||||
if text_embeddings is None and conditional_embeddings is None:
|
||||
raise ValueError("Either text_embeddings or conditional_embeddings must be specified")
|
||||
if text_embeddings is None and unconditional_embeddings is not None:
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
unconditional_embeddings, # negative embedding
|
||||
conditional_embeddings, # positive embedding
|
||||
latents.shape[0], # batch size
|
||||
)
|
||||
elif text_embeddings is None and conditional_embeddings is not None:
|
||||
# not doing cfg
|
||||
text_embeddings = conditional_embeddings
|
||||
|
||||
# CFG is comparing neg and positive, if we have concatenated embeddings
|
||||
# then we are doing it, otherwise we are not and takes half the time.
|
||||
do_classifier_free_guidance = True
|
||||
|
||||
# check if batch size of embeddings matches batch size of latents
|
||||
if latents.shape[0] == text_embeddings.text_embeds.shape[0]:
|
||||
do_classifier_free_guidance = False
|
||||
elif latents.shape[0] * 2 != text_embeddings.text_embeds.shape[0]:
|
||||
raise ValueError("Batch size of latents must be the same or half the batch size of text embeddings")
|
||||
|
||||
if self.is_xl:
|
||||
if add_time_ids is None:
|
||||
add_time_ids = self.get_time_ids_from_latents(latents)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# todo check this with larget batches
|
||||
train_util.concat_embeddings(
|
||||
add_time_ids, add_time_ids, 1
|
||||
)
|
||||
else:
|
||||
# concat to fit batch size
|
||||
add_time_ids = torch.cat([add_time_ids] * latents.shape[0])
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
|
||||
latent_model_input = self.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
added_cond_kwargs = {
|
||||
"text_embeds": text_embeddings.pooled_embeds,
|
||||
"time_ids": add_time_ids,
|
||||
}
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
).sample
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# perform guidance
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
# https://github.com/huggingface/diffusers/blob/7a91ea6c2b53f94da930a61ed571364022b21044/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py#L775
|
||||
if guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
|
||||
|
||||
else:
|
||||
if do_classifier_free_guidance:
|
||||
# if we are doing classifier free guidance, need to double up
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = self.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
).sample
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# perform guidance
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
return noise_pred
|
||||
|
||||
# ref: https://github.com/huggingface/diffusers/blob/0bab447670f47c28df60fbd2f6a0f833f75a16f5/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L746
|
||||
def diffuse_some_steps(
|
||||
self,
|
||||
latents: torch.FloatTensor,
|
||||
text_embeddings: PromptEmbeds,
|
||||
total_timesteps: int = 1000,
|
||||
start_timesteps=0,
|
||||
guidance_scale=1,
|
||||
add_time_ids=None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
for timestep in tqdm(self.noise_scheduler.timesteps[start_timesteps:total_timesteps], leave=False):
|
||||
noise_pred = self.predict_noise(
|
||||
latents,
|
||||
text_embeddings,
|
||||
timestep,
|
||||
guidance_scale=guidance_scale,
|
||||
add_time_ids=add_time_ids,
|
||||
**kwargs,
|
||||
)
|
||||
latents = self.noise_scheduler.step(noise_pred, timestep, latents).prev_sample
|
||||
|
||||
# return latents_steps
|
||||
return latents
|
||||
|
||||
def encode_prompt(self, prompt, num_images_per_prompt=1) -> PromptEmbeds:
|
||||
prompt = prompt
|
||||
@@ -86,18 +625,89 @@ class StableDiffusion:
|
||||
)
|
||||
)
|
||||
|
||||
def encode_images(
|
||||
self,
|
||||
image_list: List[torch.Tensor],
|
||||
device=None,
|
||||
dtype=None
|
||||
):
|
||||
if device is None:
|
||||
device = self.device
|
||||
if dtype is None:
|
||||
dtype = self.torch_dtype
|
||||
|
||||
latent_list = []
|
||||
# Move to vae to device if on cpu
|
||||
if self.vae.device == 'cpu':
|
||||
self.vae.to(self.device)
|
||||
# move to device and dtype
|
||||
image_list = [image.to(self.device, dtype=self.torch_dtype) for image in image_list]
|
||||
|
||||
# resize images if not divisible by 8
|
||||
for i in range(len(image_list)):
|
||||
image = image_list[i]
|
||||
if image.shape[1] % 8 != 0 or image.shape[2] % 8 != 0:
|
||||
image_list[i] = Resize((image.shape[1] // 8 * 8, image.shape[2] // 8 * 8))(image)
|
||||
|
||||
images = torch.stack(image_list)
|
||||
latents = self.vae.encode(images).latent_dist.sample()
|
||||
latents = latents * 0.18215
|
||||
latents = latents.to(device, dtype=dtype)
|
||||
|
||||
return latents
|
||||
|
||||
def encode_image_prompt_pairs(
|
||||
self,
|
||||
prompt_list: List[str],
|
||||
image_list: List[torch.Tensor],
|
||||
device=None,
|
||||
dtype=None
|
||||
):
|
||||
# todo check image types and expand and rescale as needed
|
||||
# device and dtype are for outputs
|
||||
if device is None:
|
||||
device = self.device
|
||||
if dtype is None:
|
||||
dtype = self.torch_dtype
|
||||
|
||||
embedding_list = []
|
||||
latent_list = []
|
||||
# embed the prompts
|
||||
for prompt in prompt_list:
|
||||
embedding = self.encode_prompt(prompt).to(self.device_torch, dtype=dtype)
|
||||
embedding_list.append(embedding)
|
||||
|
||||
return embedding_list, latent_list
|
||||
|
||||
def get_weight_by_name(self, name):
|
||||
# weights begin with te{te_num}_ for text encoder
|
||||
# weights begin with unet_ for unet_
|
||||
if name.startswith('te'):
|
||||
key = name[4:]
|
||||
# text encoder
|
||||
te_num = int(name[2])
|
||||
if isinstance(self.text_encoder, list):
|
||||
return self.text_encoder[te_num].state_dict()[key]
|
||||
else:
|
||||
return self.text_encoder.state_dict()[key]
|
||||
elif name.startswith('unet'):
|
||||
key = name[5:]
|
||||
# unet
|
||||
return self.unet.state_dict()[key]
|
||||
|
||||
raise ValueError(f"Unknown weight name: {name}")
|
||||
|
||||
def save(self, output_file: str, meta: OrderedDict, save_dtype=get_torch_dtype('fp16'), logit_scale=None):
|
||||
state_dict = {}
|
||||
|
||||
def update_sd(prefix, sd):
|
||||
for k, v in sd.items():
|
||||
key = prefix + k
|
||||
v = v.detach().clone()
|
||||
state_dict[key] = v.to("cpu", dtype=get_torch_dtype(save_dtype))
|
||||
|
||||
# todo see what logit scale is
|
||||
if self.is_xl:
|
||||
|
||||
state_dict = {}
|
||||
|
||||
def update_sd(prefix, sd):
|
||||
for k, v in sd.items():
|
||||
key = prefix + k
|
||||
v = v.detach().clone().to("cpu").to(get_torch_dtype(save_dtype))
|
||||
state_dict[key] = v
|
||||
|
||||
# Convert the UNet model
|
||||
update_sd("model.diffusion_model.", self.unet.state_dict())
|
||||
|
||||
@@ -107,19 +717,27 @@ class StableDiffusion:
|
||||
text_enc2_dict = convert_text_encoder_2_state_dict_to_sdxl(self.text_encoder[1].state_dict(), logit_scale)
|
||||
update_sd("conditioner.embedders.1.model.", text_enc2_dict)
|
||||
|
||||
else:
|
||||
# Convert the UNet model
|
||||
unet_state_dict = convert_unet_state_dict_to_sd(self.is_v2, self.unet.state_dict())
|
||||
update_sd("model.diffusion_model.", unet_state_dict)
|
||||
|
||||
# Convert the text encoder model
|
||||
if self.is_v2:
|
||||
make_dummy = True
|
||||
text_enc_dict = convert_text_encoder_state_dict_to_sd_v2(self.text_encoder.state_dict(), make_dummy)
|
||||
update_sd("cond_stage_model.model.", text_enc_dict)
|
||||
else:
|
||||
text_enc_dict = self.text_encoder.state_dict()
|
||||
update_sd("cond_stage_model.transformer.", text_enc_dict)
|
||||
|
||||
# Convert the VAE
|
||||
if self.vae is not None:
|
||||
vae_dict = model_util.convert_vae_state_dict(self.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}
|
||||
|
||||
if model_util.is_safetensors(output_file):
|
||||
save_file(state_dict, output_file)
|
||||
else:
|
||||
torch.save(new_ckpt, output_file, meta)
|
||||
|
||||
return key_count
|
||||
else:
|
||||
raise NotImplementedError("sdv1.x, sdv2.x is not implemented yet")
|
||||
# prepare metadata
|
||||
meta = get_meta_for_safetensors(meta)
|
||||
# make sure parent folder exists
|
||||
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
||||
save_file(state_dict, output_file, metadata=meta)
|
||||
|
||||
@@ -4,6 +4,10 @@ import json
|
||||
import os
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
import sys
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from diffusers import (
|
||||
StableDiffusionPipeline,
|
||||
@@ -30,13 +34,16 @@ SCHEDLER_SCHEDULE = "scaled_linear"
|
||||
|
||||
|
||||
def get_torch_dtype(dtype_str):
|
||||
# if it is a torch dtype, return it
|
||||
if isinstance(dtype_str, torch.dtype):
|
||||
return dtype_str
|
||||
if dtype_str == "float" or dtype_str == "fp32" or dtype_str == "single" or dtype_str == "float32":
|
||||
return torch.float
|
||||
if dtype_str == "fp16" or dtype_str == "half" or dtype_str == "float16":
|
||||
return torch.float16
|
||||
if dtype_str == "bf16" or dtype_str == "bfloat16":
|
||||
return torch.bfloat16
|
||||
return None
|
||||
return dtype_str
|
||||
|
||||
|
||||
def replace_filewords_prompt(prompt, args: argparse.Namespace):
|
||||
|
||||
Reference in New Issue
Block a user