Compare commits
41 Commits
sdxl
...
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 | ||
|
|
8b8d53888d | ||
|
|
63cacf4362 | ||
|
|
c1b1e800df | ||
|
|
7726911562 | ||
|
|
c01673f1b5 | ||
|
|
c35b78f0d4 | ||
|
|
8ba1b11557 | ||
|
|
1e50b39442 | ||
|
|
5fc2bb5d9c | ||
|
|
c7640b0865 | ||
|
|
b2e2e4bf47 | ||
|
|
596e57a6a6 | ||
|
|
6ab8b8b0f1 |
4
.gitignore
vendored
4
.gitignore
vendored
@@ -170,4 +170,6 @@ cython_debug/
|
|||||||
!/config/examples
|
!/config/examples
|
||||||
!/config/_PUT_YOUR_CONFIGS_HERE).txt
|
!/config/_PUT_YOUR_CONFIGS_HERE).txt
|
||||||
/output/*
|
/output/*
|
||||||
!/output/.gitkeep
|
!/output/.gitkeep
|
||||||
|
/extensions/*
|
||||||
|
!/extensions/example
|
||||||
105
README.md
105
README.md
@@ -29,7 +29,9 @@ cd ai-toolkit
|
|||||||
git submodule update --init --recursive
|
git submodule update --init --recursive
|
||||||
python3 -m venv venv
|
python3 -m venv venv
|
||||||
source venv/bin/activate
|
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
|
pip3 install -r requirements.txt
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -40,6 +42,18 @@ pip3 install -r requirements.txt
|
|||||||
I have so many hodge podge scripts I am going to be moving over to this that I use in my ML work. But this is what is
|
I have so many hodge podge scripts I am going to be moving over to this that I use in my ML work. But this is what is
|
||||||
here so far.
|
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
|
### 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
|
It is based on the extractor in the [LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS) tool, but adding some QOL features
|
||||||
@@ -64,9 +78,38 @@ Most people used fixed, which is traditional fixed dimension extraction.
|
|||||||
|
|
||||||
`process` is an array of different processes to run. You can add a few and mix and match. One LoRA, one LyCON, etc.
|
`process` is an array of different processes to run. You can add a few and mix and match. One LoRA, one LyCON, etc.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### LoRA Rescale
|
||||||
|
|
||||||
|
Change `<lora:my_lora:4.6>` to `<lora:my_lora:1.0>` or whatever you want with the same effect.
|
||||||
|
A tool for rescaling a LoRA's weights. Should would with LoCON as well, but I have not tested it.
|
||||||
|
It all runs off a config file, which you can find an example of in `config/examples/mod_lora_scale.yml`.
|
||||||
|
Just copy that file, into the `config` folder, and rename it to `whatever_you_want.yml`.
|
||||||
|
Then you can edit the file to your liking. and call it like so:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python3 run.py config/whatever_you_want.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
You can also put a full path to a config file, if you want to keep it somewhere else.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python3 run.py "/home/user/whatever_you_want.yml"
|
||||||
|
```
|
||||||
|
|
||||||
|
More notes on how it works are available in the example config file itself. This is useful when making
|
||||||
|
all LoRAs, as the ideal weight is rarely 1.0, but now you can fix that. For sliders, they can have weird scales form -2 to 2
|
||||||
|
or even -15 to 15. This will allow you to dile it in so they all have your desired scale
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
### LoRA Slider Trainer
|
### 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).
|
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)
|
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
|
But has been heavily modified to create sliders rather than erasing concepts. I have a lot more plans on this, but it is
|
||||||
@@ -89,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
|
## WIP Tools
|
||||||
|
|
||||||
|
|
||||||
@@ -108,14 +168,53 @@ Just went in and out. It is much worse on smaller faces than shown here.
|
|||||||
|
|
||||||
## TODO
|
## TODO
|
||||||
- [X] Add proper regs on sliders
|
- [X] Add proper regs on sliders
|
||||||
- [ ] Add SDXL support (base model only for now)
|
- [X] Add SDXL support (base model only for now)
|
||||||
- [ ] Add plain erasing
|
- [ ] Add plain erasing
|
||||||
- [ ] Make Textual inversion network trainer (network that spits out TI embeddings)
|
- [ ] Make Textual inversion network trainer (network that spits out TI embeddings)
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Change Log
|
## Change Log
|
||||||
#### 2021-07-30
|
|
||||||
|
#### 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.
|
||||||
|
|
||||||
|
Unfortunately, I am too lazy to write a proper changelog with all the changes.
|
||||||
|
|
||||||
|
I added SDXL training to sliders... but.. it does not work properly.
|
||||||
|
The slider training relies on a model's ability to understand that an unconditional (negative prompt)
|
||||||
|
means you do not want that concept in the output. SDXL does not understand this for whatever reason,
|
||||||
|
which makes separating out
|
||||||
|
concepts within the model hard. I am sure the community will find a way to fix this
|
||||||
|
over time, but for now, it is not
|
||||||
|
going to work properly. And if any of you are thinking "Could we maybe fix it by adding 1 or 2 more text
|
||||||
|
encoders to the model as well as a few more entirely separate diffusion networks?" No. God no. It just needs a little
|
||||||
|
training without every experimental new paper added to it. The KISS principal.
|
||||||
|
|
||||||
|
|
||||||
|
#### 2023-07-30
|
||||||
Added "anchors" to the slider trainer. This allows you to set a prompt that will be used as a
|
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
|
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
|
||||||
48
config/examples/mod_lora_scale.yaml
Normal file
48
config/examples/mod_lora_scale.yaml
Normal file
@@ -0,0 +1,48 @@
|
|||||||
|
---
|
||||||
|
job: mod
|
||||||
|
config:
|
||||||
|
name: name_of_your_model_v1
|
||||||
|
process:
|
||||||
|
- type: rescale_lora
|
||||||
|
# path to your current lora model
|
||||||
|
input_path: "/path/to/lora/lora.safetensors"
|
||||||
|
# output path for your new lora model, can be the same as input_path to replace
|
||||||
|
output_path: "/path/to/lora/output_lora_v1.safetensors"
|
||||||
|
# replaces meta with the meta below (plus minimum meta fields)
|
||||||
|
# if false, we will leave the meta alone except for updating hashes (sd-script hashes)
|
||||||
|
replace_meta: true
|
||||||
|
# how to adjust, we can scale the up_down weights or the alpha
|
||||||
|
# up_down is the default and probably the best, they will both net the same outputs
|
||||||
|
# would only affect rare NaN cases and maybe merging with old merge tools
|
||||||
|
scale_target: 'up_down'
|
||||||
|
# precision to save, fp16 is the default and standard
|
||||||
|
save_dtype: fp16
|
||||||
|
# current_weight is the ideal weight you use as a multiplier when using the lora
|
||||||
|
# IE in automatic1111 <lora:my_lora:6.0> the 6.0 is the current_weight
|
||||||
|
# you can do negatives here too if you want to flip the lora
|
||||||
|
current_weight: 6.0
|
||||||
|
# target_weight is the ideal weight you use as a multiplier when using the lora
|
||||||
|
# instead of the one above. IE in automatic1111 instead of using <lora:my_lora:6.0>
|
||||||
|
# we want to use <lora:my_lora:1.0> so 1.0 is the target_weight
|
||||||
|
target_weight: 1.0
|
||||||
|
|
||||||
|
# base model for the lora
|
||||||
|
# this is just used to add meta so automatic111 knows which model it is for
|
||||||
|
# assume v1.5 if these are not set
|
||||||
|
is_xl: false
|
||||||
|
is_v2: false
|
||||||
|
meta:
|
||||||
|
# this is only used if you set replace_meta to true above
|
||||||
|
name: "[name]" # [name] gets replaced with the name above
|
||||||
|
description: A short description of your lora
|
||||||
|
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.
|
||||||
@@ -7,7 +7,7 @@ job: train
|
|||||||
config:
|
config:
|
||||||
# the name will be used to create a folder in the output folder
|
# 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
|
# 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
|
# folder will be created with name above in folder below
|
||||||
# it can be relative to the project root or absolute
|
# it can be relative to the project root or absolute
|
||||||
training_folder: "output/LoRA"
|
training_folder: "output/LoRA"
|
||||||
@@ -24,7 +24,7 @@ config:
|
|||||||
type: "lierla"
|
type: "lierla"
|
||||||
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
|
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
|
||||||
rank: 8
|
rank: 8
|
||||||
alpha: 1.0 # just leave it
|
alpha: 4 # Do about half of rank
|
||||||
|
|
||||||
# training config
|
# training config
|
||||||
train:
|
train:
|
||||||
@@ -33,7 +33,9 @@ config:
|
|||||||
# how many steps to train. More is not always better. I rarely go over 1000
|
# how many steps to train. More is not always better. I rarely go over 1000
|
||||||
steps: 500
|
steps: 500
|
||||||
# I have had good results with 4e-4 to 1e-4 at 500 steps
|
# 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 the unet. I recommend leaving this true
|
||||||
train_unet: true
|
train_unet: true
|
||||||
# train the text encoder. I don't recommend this unless you have a special use case
|
# 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)
|
# not the description of it (text encoder)
|
||||||
train_text_encoder: false
|
train_text_encoder: false
|
||||||
|
|
||||||
|
|
||||||
# just leave unless you know what you are doing
|
# just leave unless you know what you are doing
|
||||||
# also supports "dadaptation" but set lr to 1 if you use that,
|
# also supports "dadaptation" but set lr to 1 if you use that,
|
||||||
# but it learns too fast and I don't recommend it
|
# but it learns too fast and I don't recommend it
|
||||||
@@ -51,11 +54,13 @@ config:
|
|||||||
# while training. Just leave it
|
# while training. Just leave it
|
||||||
max_denoising_steps: 40
|
max_denoising_steps: 40
|
||||||
# works great at 1. I do 1 even with my 4090.
|
# 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
|
batch_size: 1
|
||||||
# bf16 works best if your GPU supports it (modern)
|
# bf16 works best if your GPU supports it (modern)
|
||||||
dtype: bf16 # fp32, bf16, fp16
|
dtype: bf16 # fp32, bf16, fp16
|
||||||
# if you have it, use it. It is faster and better
|
# 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
|
# 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
|
# although, the way we train sliders is comparative, so it probably won't work anyway
|
||||||
noise_offset: 0.0
|
noise_offset: 0.0
|
||||||
@@ -66,11 +71,17 @@ config:
|
|||||||
name_or_path: "runwayml/stable-diffusion-v1-5"
|
name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||||
is_v2: false # for v2 models
|
is_v2: false # for v2 models
|
||||||
is_v_pred: false # for v-prediction models (most 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
|
# saving config
|
||||||
save:
|
save:
|
||||||
dtype: float16 # precision to save. I recommend float16
|
dtype: float16 # precision to save. I recommend float16
|
||||||
save_every: 50 # save every this many steps
|
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
|
# sampling config
|
||||||
sample:
|
sample:
|
||||||
@@ -88,21 +99,22 @@ config:
|
|||||||
# --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive
|
# --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
|
# slide are good tests. will inherit sample.network_multiplier if not set
|
||||||
# --n [string] # negative prompt, will inherit sample.neg if not set
|
# --n [string] # negative prompt, will inherit sample.neg if not set
|
||||||
|
|
||||||
# Only 75 tokens allowed currently
|
# Only 75 tokens allowed currently
|
||||||
prompts: # our example is an animal slider, neg: dog, pos: cat
|
# I like to do a wide positive and negative spread so I can see a good range and stop
|
||||||
- "a golden retriever --m -5"
|
# early if the network is braking down
|
||||||
- "a golden retriever --m -3"
|
prompts:
|
||||||
- "a golden retriever --m 3"
|
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5"
|
||||||
- "a golden retriever --m 5"
|
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3"
|
||||||
- "calico cat --m -5"
|
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3"
|
||||||
- "calico cat --m -3"
|
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5"
|
||||||
- "calico cat --m 3"
|
- "a golden retriever sitting on a leather couch, --m -5"
|
||||||
- "calico cat --m 5"
|
- "a golden retriever sitting on a leather couch --m -3"
|
||||||
- "an elephant --m -5"
|
- "a golden retriever sitting on a leather couch --m 3"
|
||||||
- "an elephant --m -3"
|
- "a golden retriever sitting on a leather couch --m 5"
|
||||||
- "an elephant --m 3"
|
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5"
|
||||||
- "an elephant --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
|
# negative prompt used on all prompts above as default if they don't have one
|
||||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
|
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
|
||||||
# seed for sampling. 42 is the answer for everything
|
# 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
|
# 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
|
# 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
|
# 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
|
# you can do as many as you want here
|
||||||
resolutions:
|
resolutions:
|
||||||
- [ 512, 512 ]
|
- [ 512, 512 ]
|
||||||
# - [ 512, 768 ]
|
# - [ 512, 768 ]
|
||||||
# - [ 768, 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,
|
# 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
|
# but they can conflict outweigh each other. Other than experimenting, I recommend
|
||||||
# just doing one for good results
|
# just doing one for good results
|
||||||
@@ -146,7 +163,9 @@ config:
|
|||||||
# a keyword necessarily but what the model understands the concept to represent.
|
# 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
|
# "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
|
# 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.
|
# 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
|
# 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
|
# 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,
|
# 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"
|
# if you want to train on fat people, you would use "an extremely fat, morbidly obese person"
|
||||||
# as the prompt. Not just "fat 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
|
# 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
|
# 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.
|
# 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
|
# 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
|
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
|
# anchors are prompts that we will try to hold on to while training the slider
|
||||||
# without directly overlapping it. For example, if you are training on a person smiling,
|
# these are NOT necessary and can prevent the slider from converging if not done right
|
||||||
# you would use "a person with a face mask" as an anchor. It is a person, the image is the same
|
# leave them off if you are having issues, but they can help lock the network
|
||||||
# regardless if they are smiling or not
|
# on certain concepts to help prevent catastrophic forgetting
|
||||||
anchors:
|
# you want these to generate an image that is not your target_class, but close to it
|
||||||
# only positive prompt for now
|
# is fine as long as it does not directly overlap it.
|
||||||
- prompt: "a woman"
|
# For example, if you are training on a person smiling,
|
||||||
neg_prompt: "animal"
|
# you could use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||||
# the multiplier applied to the LoRA when this is run.
|
# regardless if they are smiling or not, however, the closer the concept is to the target_class
|
||||||
# higher will give it more weight but also help keep the lora from collapsing
|
# the less the multiplier needs to be. Keep multipliers less than 1.0 for anchors usually
|
||||||
multiplier: 8.0
|
# for close concepts, you want to be closer to 0.1 or 0.2
|
||||||
- prompt: "a man"
|
# these will slow down training. I am leaving them off for the demo
|
||||||
neg_prompt: "animal"
|
|
||||||
multiplier: 8.0
|
# anchors:
|
||||||
- prompt: "a person"
|
# - prompt: "a woman"
|
||||||
neg_prompt: "animal"
|
# neg_prompt: "animal"
|
||||||
multiplier: 8.0
|
# # 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.
|
# 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.
|
# 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 = OrderedDict()
|
||||||
v["name"] = "ai-toolkit"
|
v["name"] = "ai-toolkit"
|
||||||
v["repo"] = "https://github.com/ostris/ai-toolkit"
|
v["repo"] = "https://github.com/ostris/ai-toolkit"
|
||||||
v["version"] = "0.0.1"
|
v["version"] = "0.0.4"
|
||||||
|
|
||||||
software_meta = v
|
software_meta = v
|
||||||
|
|||||||
@@ -60,7 +60,11 @@ class BaseJob:
|
|||||||
|
|
||||||
# check if dict key is process type
|
# check if dict key is process type
|
||||||
if process['type'] in process_dict:
|
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))
|
self.process.append(ProcessClass(i, self, process))
|
||||||
else:
|
else:
|
||||||
raise ValueError(f'config file is invalid. Unknown process type: {process["type"]}')
|
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()
|
||||||
28
jobs/ModJob.py
Normal file
28
jobs/ModJob.py
Normal file
@@ -0,0 +1,28 @@
|
|||||||
|
import os
|
||||||
|
from collections import OrderedDict
|
||||||
|
from jobs import BaseJob
|
||||||
|
from toolkit.metadata import get_meta_for_safetensors
|
||||||
|
from toolkit.train_tools import get_torch_dtype
|
||||||
|
|
||||||
|
process_dict = {
|
||||||
|
'rescale_lora': 'ModRescaleLoraProcess',
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ModJob(BaseJob):
|
||||||
|
|
||||||
|
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()
|
||||||
@@ -17,8 +17,11 @@ sys.path.append(REPOS_ROOT)
|
|||||||
process_dict = {
|
process_dict = {
|
||||||
'vae': 'TrainVAEProcess',
|
'vae': 'TrainVAEProcess',
|
||||||
'slider': 'TrainSliderProcess',
|
'slider': 'TrainSliderProcess',
|
||||||
|
'slider_old': 'TrainSliderProcessOld',
|
||||||
'lora_hack': 'TrainLoRAHack',
|
'lora_hack': 'TrainLoRAHack',
|
||||||
'rescale_sd': 'TrainSDRescaleProcess',
|
'rescale_sd': 'TrainSDRescaleProcess',
|
||||||
|
'esrgan': 'TrainESRGANProcess',
|
||||||
|
'reference': 'TrainReferenceProcess',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -34,18 +37,9 @@ class TrainJob(BaseJob):
|
|||||||
# self.mixed_precision = self.get_conf('mixed_precision', False) # fp16
|
# self.mixed_precision = self.get_conf('mixed_precision', False) # fp16
|
||||||
self.log_dir = self.get_conf('log_dir', None)
|
self.log_dir = self.get_conf('log_dir', None)
|
||||||
|
|
||||||
self.writer = None
|
|
||||||
self.setup_tensorboard()
|
|
||||||
|
|
||||||
# loads the processes from the config
|
# loads the processes from the config
|
||||||
self.load_processes(process_dict)
|
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):
|
def run(self):
|
||||||
super().run()
|
super().run()
|
||||||
@@ -54,12 +48,3 @@ class TrainJob(BaseJob):
|
|||||||
|
|
||||||
for process in self.process:
|
for process in self.process:
|
||||||
process.run()
|
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)
|
|
||||||
|
|||||||
@@ -2,3 +2,6 @@ from .BaseJob import BaseJob
|
|||||||
from .ExtractJob import ExtractJob
|
from .ExtractJob import ExtractJob
|
||||||
from .TrainJob import TrainJob
|
from .TrainJob import TrainJob
|
||||||
from .MergeJob import MergeJob
|
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 copy
|
||||||
import json
|
import json
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import ForwardRef
|
|
||||||
|
|
||||||
|
|
||||||
class BaseProcess:
|
class BaseProcess(object):
|
||||||
meta: OrderedDict
|
meta: OrderedDict
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -16,6 +15,8 @@ class BaseProcess:
|
|||||||
self.process_id = process_id
|
self.process_id = process_id
|
||||||
self.job = job
|
self.job = job
|
||||||
self.config = config
|
self.config = config
|
||||||
|
self.raw_process_config = config
|
||||||
|
self.name = self.get_conf('name', self.job.name)
|
||||||
self.meta = copy.deepcopy(self.job.meta)
|
self.meta = copy.deepcopy(self.job.meta)
|
||||||
print(json.dumps(self.config, indent=4))
|
print(json.dumps(self.config, indent=4))
|
||||||
|
|
||||||
|
|||||||
@@ -1,34 +1,26 @@
|
|||||||
import glob
|
import glob
|
||||||
import time
|
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
import os
|
import os
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
import diffusers
|
from torch.utils.data import DataLoader
|
||||||
from safetensors import safe_open
|
|
||||||
|
|
||||||
from library import sdxl_train_util, sdxl_model_util
|
|
||||||
from toolkit.kohya_model_util import load_vae
|
|
||||||
from toolkit.lora_special import LoRASpecialNetwork
|
from toolkit.lora_special import LoRASpecialNetwork
|
||||||
from toolkit.optimizer import get_optimizer
|
from toolkit.optimizer import get_optimizer
|
||||||
from toolkit.paths import REPOS_ROOT
|
|
||||||
import sys
|
|
||||||
|
|
||||||
sys.path.append(REPOS_ROOT)
|
from toolkit.scheduler import get_lr_scheduler
|
||||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
from toolkit.stable_diffusion_model import StableDiffusion
|
||||||
|
|
||||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, DDPMScheduler
|
|
||||||
|
|
||||||
from jobs.process import BaseTrainProcess
|
from jobs.process import BaseTrainProcess
|
||||||
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors
|
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 gc
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from leco import train_util, model_util
|
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \
|
||||||
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig
|
GenerateImageConfig
|
||||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
|
||||||
|
|
||||||
|
|
||||||
def flush():
|
def flush():
|
||||||
@@ -36,212 +28,104 @@ def flush():
|
|||||||
gc.collect()
|
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):
|
class BaseSDTrainProcess(BaseTrainProcess):
|
||||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
sd: StableDiffusion
|
||||||
|
|
||||||
|
def __init__(self, process_id: int, job, config: OrderedDict, custom_pipeline=None):
|
||||||
super().__init__(process_id, job, config)
|
super().__init__(process_id, job, config)
|
||||||
|
self.custom_pipeline = custom_pipeline
|
||||||
self.step_num = 0
|
self.step_num = 0
|
||||||
self.start_step = 0
|
self.start_step = 0
|
||||||
self.device = self.get_conf('device', self.job.device)
|
self.device = self.get_conf('device', self.job.device)
|
||||||
self.device_torch = torch.device(self.device)
|
self.device_torch = torch.device(self.device)
|
||||||
self.network_config = NetworkConfig(**self.get_conf('network', None))
|
network_config = self.get_conf('network', None)
|
||||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
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.train_config = TrainConfig(**self.get_conf('train', {}))
|
||||||
self.model_config = ModelConfig(**self.get_conf('model', {}))
|
self.model_config = ModelConfig(**self.get_conf('model', {}))
|
||||||
self.save_config = SaveConfig(**self.get_conf('save', {}))
|
self.save_config = SaveConfig(**self.get_conf('save', {}))
|
||||||
self.sample_config = SampleConfig(**self.get_conf('sample', {}))
|
self.sample_config = SampleConfig(**self.get_conf('sample', {}))
|
||||||
self.first_sample_config = SampleConfig(
|
first_sample_config = self.get_conf('first_sample', None)
|
||||||
**self.get_conf('first_sample', {})) if 'first_sample' in self.config else self.sample_config
|
if first_sample_config is not None:
|
||||||
|
self.has_first_sample_requested = True
|
||||||
|
self.first_sample_config = SampleConfig(**first_sample_config)
|
||||||
|
else:
|
||||||
|
self.has_first_sample_requested = False
|
||||||
|
self.first_sample_config = self.sample_config
|
||||||
self.logging_config = LogingConfig(**self.get_conf('logging', {}))
|
self.logging_config = LogingConfig(**self.get_conf('logging', {}))
|
||||||
self.optimizer = None
|
self.optimizer = None
|
||||||
self.lr_scheduler = None
|
self.lr_scheduler = None
|
||||||
self.sd: 'StableDiffusion' = None
|
self.data_loader: Union[DataLoader, None] = None
|
||||||
|
|
||||||
# sdxl stuff
|
self.sd = StableDiffusion(
|
||||||
self.logit_scale = None
|
device=self.device,
|
||||||
self.ckppt_info = None
|
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
|
self.network = None
|
||||||
|
|
||||||
def sample(self, step=None, is_first=False):
|
def sample(self, step=None, is_first=False):
|
||||||
sample_folder = os.path.join(self.save_root, 'samples')
|
sample_folder = os.path.join(self.save_root, 'samples')
|
||||||
if not os.path.exists(sample_folder):
|
gen_img_config_list = []
|
||||||
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)
|
|
||||||
|
|
||||||
sample_config = self.first_sample_config if is_first else self.sample_config
|
sample_config = self.first_sample_config if is_first else self.sample_config
|
||||||
|
|
||||||
start_seed = sample_config.seed
|
start_seed = sample_config.seed
|
||||||
start_multiplier = self.network.multiplier
|
|
||||||
current_seed = start_seed
|
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)
|
step_num = ''
|
||||||
with self.network:
|
if step is not None:
|
||||||
with torch.no_grad():
|
# zero-pad 9 digits
|
||||||
if self.network is not None:
|
step_num = f"_{str(step).zfill(9)}"
|
||||||
assert self.network.is_active
|
|
||||||
if self.logging_config.verbose:
|
|
||||||
print("network_state", {
|
|
||||||
'is_active': self.network.is_active,
|
|
||||||
'multiplier': self.network.multiplier,
|
|
||||||
})
|
|
||||||
|
|
||||||
for i in tqdm(range(len(sample_config.prompts)), desc=f"Generating Samples - step: {step}",
|
filename = f"[time]_{step_num}_[count].png"
|
||||||
leave=False):
|
|
||||||
raw_prompt = sample_config.prompts[i]
|
|
||||||
|
|
||||||
neg = sample_config.neg
|
output_path = os.path.join(sample_folder, filename)
|
||||||
multiplier = sample_config.network_multiplier
|
|
||||||
p_split = raw_prompt.split('--')
|
|
||||||
prompt = p_split[0].strip()
|
|
||||||
height = sample_config.height
|
|
||||||
width = sample_config.width
|
|
||||||
|
|
||||||
if len(p_split) > 1:
|
gen_img_config_list.append(GenerateImageConfig(
|
||||||
for split in p_split:
|
prompt=sample_config.prompts[i], # it will autoparse the prompt
|
||||||
flag = split[:1]
|
width=sample_config.width,
|
||||||
content = split[1:].strip()
|
height=sample_config.height,
|
||||||
if flag == 'n':
|
negative_prompt=sample_config.neg,
|
||||||
neg = content
|
seed=current_seed,
|
||||||
elif flag == 'm':
|
guidance_scale=sample_config.guidance_scale,
|
||||||
# multiplier
|
guidance_rescale=sample_config.guidance_rescale,
|
||||||
multiplier = float(content)
|
num_inference_steps=sample_config.sample_steps,
|
||||||
elif flag == 'w':
|
network_multiplier=sample_config.network_multiplier,
|
||||||
# multiplier
|
output_path=output_path,
|
||||||
width = int(content)
|
))
|
||||||
elif flag == 'h':
|
|
||||||
# multiplier
|
|
||||||
height = int(content)
|
|
||||||
|
|
||||||
height = max(64, height - height % 8) # round to divisible by 8
|
# send to be generated
|
||||||
width = max(64, width - width % 8) # round to divisible by 8
|
self.sd.generate_images(gen_img_config_list)
|
||||||
|
|
||||||
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,
|
|
||||||
).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'])
|
|
||||||
|
|
||||||
def update_training_metadata(self):
|
def update_training_metadata(self):
|
||||||
dict = OrderedDict({
|
o_dict = OrderedDict({
|
||||||
"training_info": self.get_training_info()
|
"training_info": self.get_training_info()
|
||||||
})
|
})
|
||||||
if self.model_config.is_v2:
|
if self.model_config.is_v2:
|
||||||
dict['ss_v2'] = True
|
o_dict['ss_v2'] = True
|
||||||
|
o_dict['ss_base_model_version'] = 'sd_2.1'
|
||||||
|
|
||||||
if self.model_config.is_xl:
|
elif self.model_config.is_xl:
|
||||||
dict['ss_base_model_version'] = 'sdxl_1.0'
|
o_dict['ss_base_model_version'] = 'sdxl_1.0'
|
||||||
|
else:
|
||||||
|
o_dict['ss_base_model_version'] = 'sd_1.5'
|
||||||
|
|
||||||
dict['ss_output_name'] = self.job.name
|
o_dict = add_base_model_info_to_meta(
|
||||||
|
o_dict,
|
||||||
|
is_v2=self.model_config.is_v2,
|
||||||
|
is_xl=self.model_config.is_xl,
|
||||||
|
)
|
||||||
|
o_dict['ss_output_name'] = self.job.name
|
||||||
|
|
||||||
self.add_meta(dict)
|
self.add_meta(o_dict)
|
||||||
|
|
||||||
def get_training_info(self):
|
def get_training_info(self):
|
||||||
info = OrderedDict({
|
info = OrderedDict({
|
||||||
@@ -299,6 +183,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.print(f"Saved to {file_path}")
|
self.print(f"Saved to {file_path}")
|
||||||
|
self.clean_up_saves()
|
||||||
|
|
||||||
# Called before the model is loaded
|
# Called before the model is loaded
|
||||||
def hook_before_model_load(self):
|
def hook_before_model_load(self):
|
||||||
@@ -312,153 +197,10 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
def hook_before_train_loop(self):
|
def hook_before_train_loop(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def get_latent_noise(
|
def hook_train_loop(self, batch=None):
|
||||||
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")
|
|
||||||
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):
|
|
||||||
# return loss
|
# return loss
|
||||||
return 0.0
|
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.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)
|
|
||||||
# todo LECOs code looks like it is omitting noise_pred
|
|
||||||
# noise_pred = train_util.predict_noise_xl(
|
|
||||||
# self.sd.unet,
|
|
||||||
# self.sd.noise_scheduler,
|
|
||||||
# timestep,
|
|
||||||
# latents,
|
|
||||||
# text_embeddings.text_embeds,
|
|
||||||
# text_embeddings.pooled_embeds,
|
|
||||||
# add_time_ids,
|
|
||||||
# guidance_scale=guidance_scale,
|
|
||||||
# guidance_rescale=guidance_rescale
|
|
||||||
# )
|
|
||||||
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)
|
|
||||||
guided_target = 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
|
|
||||||
# noise_pred = rescale_noise_cfg(
|
|
||||||
# noise_pred, noise_pred_text, guidance_rescale=guidance_rescale
|
|
||||||
# )
|
|
||||||
|
|
||||||
noise_pred = guided_target
|
|
||||||
|
|
||||||
else:
|
|
||||||
noise_pred = train_util.predict_noise(
|
|
||||||
self.sd.unet,
|
|
||||||
self.sd.noise_scheduler,
|
|
||||||
timestep,
|
|
||||||
latents,
|
|
||||||
text_embeddings.text_embeds if hasattr(text_embeddings, 'text_embeds') else text_embeddings,
|
|
||||||
guidance_scale=guidance_scale
|
|
||||||
)
|
|
||||||
|
|
||||||
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):
|
def get_latest_save_path(self):
|
||||||
# get latest saved step
|
# get latest saved step
|
||||||
if os.path.exists(self.save_root):
|
if os.path.exists(self.save_root):
|
||||||
@@ -486,84 +228,56 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
print("load_weights not implemented for non-network models")
|
print("load_weights not implemented for non-network models")
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
super().run()
|
# run base process run
|
||||||
|
BaseTrainProcess.run(self)
|
||||||
### HOOK ###
|
### HOOK ###
|
||||||
self.hook_before_model_load()
|
self.hook_before_model_load()
|
||||||
|
# run base sd process run
|
||||||
|
self.sd.load_model()
|
||||||
|
|
||||||
dtype = get_torch_dtype(self.train_config.dtype)
|
dtype = get_torch_dtype(self.train_config.dtype)
|
||||||
|
|
||||||
if self.model_config.is_xl:
|
# 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
|
||||||
|
|
||||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
|
||||||
self.model_config.name_or_path,
|
|
||||||
dtype=dtype,
|
|
||||||
scheduler_type='pndm',
|
|
||||||
device=self.device_torch
|
|
||||||
)
|
|
||||||
text_encoders = [pipe.text_encoder, pipe.text_encoder_2]
|
|
||||||
tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
|
||||||
unet = pipe.unet
|
|
||||||
noise_scheduler = pipe.scheduler
|
|
||||||
vae = pipe.vae.to('cpu', dtype=dtype)
|
|
||||||
vae.eval()
|
|
||||||
vae.set_use_memory_efficient_attention_xformers(True)
|
|
||||||
|
|
||||||
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
|
|
||||||
tokenizer = tokenizer
|
|
||||||
del pipe
|
|
||||||
flush()
|
|
||||||
|
|
||||||
|
|
||||||
else:
|
|
||||||
tokenizer, text_encoder, unet, noise_scheduler = model_util.load_models(
|
|
||||||
self.model_config.name_or_path,
|
|
||||||
scheduler_name=self.train_config.noise_scheduler,
|
|
||||||
v2=self.model_config.is_v2,
|
|
||||||
v_pred=self.model_config.is_v_pred,
|
|
||||||
)
|
|
||||||
|
|
||||||
text_encoder.to(self.device_torch, dtype=dtype)
|
|
||||||
text_encoder.eval()
|
|
||||||
vae = load_vae(self.model_config.name_or_path, dtype=dtype).to('cpu', dtype=dtype)
|
|
||||||
vae.eval()
|
|
||||||
flush()
|
|
||||||
|
|
||||||
|
|
||||||
# just for now or of we want to load a custom one
|
|
||||||
# put on cpu for now, we only need it when sampling
|
|
||||||
# vae = load_vae(self.model_config.name_or_path, dtype=dtype).to('cpu', dtype=dtype)
|
|
||||||
# vae.eval()
|
|
||||||
self.sd = StableDiffusion(vae, tokenizer, text_encoder, unet, noise_scheduler, is_xl=self.model_config.is_xl)
|
|
||||||
|
|
||||||
unet.to(self.device_torch, dtype=dtype)
|
|
||||||
if self.train_config.xformers:
|
if self.train_config.xformers:
|
||||||
|
vae.set_use_memory_efficient_attention_xformers(True)
|
||||||
unet.enable_xformers_memory_efficient_attention()
|
unet.enable_xformers_memory_efficient_attention()
|
||||||
if self.train_config.gradient_checkpointing:
|
if self.train_config.gradient_checkpointing:
|
||||||
unet.enable_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.requires_grad_(False)
|
||||||
unet.eval()
|
unet.eval()
|
||||||
|
vae = vae.to(torch.device('cpu'), dtype=dtype)
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
vae.eval()
|
||||||
|
|
||||||
if self.network_config is not None:
|
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(
|
self.network = LoRASpecialNetwork(
|
||||||
text_encoder=text_encoder,
|
text_encoder=text_encoder,
|
||||||
unet=unet,
|
unet=unet,
|
||||||
lora_dim=self.network_config.linear,
|
lora_dim=self.network_config.linear,
|
||||||
multiplier=1.0,
|
multiplier=1.0,
|
||||||
alpha=self.network_config.alpha,
|
alpha=self.network_config.linear_alpha,
|
||||||
train_unet=self.train_config.train_unet,
|
train_unet=self.train_config.train_unet,
|
||||||
train_text_encoder=self.train_config.train_text_encoder,
|
train_text_encoder=self.train_config.train_text_encoder,
|
||||||
conv_lora_dim=conv,
|
conv_lora_dim=self.network_config.conv,
|
||||||
conv_alpha=self.network_config.alpha if conv is not None else None,
|
conv_alpha=self.network_config.conv_alpha,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.network.force_to(self.device_torch, dtype=dtype)
|
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(
|
self.network.apply_to(
|
||||||
text_encoder,
|
text_encoder,
|
||||||
@@ -580,6 +294,9 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
default_lr=self.train_config.lr
|
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()
|
latest_save_path = self.get_latest_save_path()
|
||||||
if latest_save_path is not None:
|
if latest_save_path is not None:
|
||||||
self.print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
|
self.print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
|
||||||
@@ -588,14 +305,19 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
self.network.multiplier = 1.0
|
self.network.multiplier = 1.0
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
params = []
|
params = []
|
||||||
# assume dreambooth/finetune
|
# assume dreambooth/finetune
|
||||||
if self.train_config.train_text_encoder:
|
if self.train_config.train_text_encoder:
|
||||||
text_encoder.requires_grad_(True)
|
if self.sd.is_xl:
|
||||||
text_encoder.train()
|
for te in text_encoder:
|
||||||
params += text_encoder.parameters()
|
te.requires_grad_(True)
|
||||||
|
te.train()
|
||||||
|
params += te.parameters()
|
||||||
|
else:
|
||||||
|
text_encoder.requires_grad_(True)
|
||||||
|
text_encoder.train()
|
||||||
|
params += text_encoder.parameters()
|
||||||
if self.train_config.train_unet:
|
if self.train_config.train_unet:
|
||||||
unet.requires_grad_(True)
|
unet.requires_grad_(True)
|
||||||
unet.train()
|
unet.train()
|
||||||
@@ -609,11 +331,11 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
optimizer_params=self.train_config.optimizer_params)
|
optimizer_params=self.train_config.optimizer_params)
|
||||||
self.optimizer = optimizer
|
self.optimizer = optimizer
|
||||||
|
|
||||||
lr_scheduler = train_util.get_lr_scheduler(
|
lr_scheduler = get_lr_scheduler(
|
||||||
self.train_config.lr_scheduler,
|
self.train_config.lr_scheduler,
|
||||||
optimizer,
|
optimizer,
|
||||||
max_iterations=self.train_config.steps,
|
max_iterations=self.train_config.steps,
|
||||||
lr_min=self.train_config.lr / 100, # not sure why leco did this, but ill do it to
|
lr_min=self.train_config.lr / 100,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.lr_scheduler = lr_scheduler
|
self.lr_scheduler = lr_scheduler
|
||||||
@@ -621,28 +343,51 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
|||||||
### HOOK ###
|
### HOOK ###
|
||||||
self.hook_before_train_loop()
|
self.hook_before_train_loop()
|
||||||
|
|
||||||
|
if self.has_first_sample_requested:
|
||||||
|
self.print("Generating first sample from first sample config")
|
||||||
|
self.sample(0, is_first=True)
|
||||||
|
|
||||||
# sample first
|
# sample first
|
||||||
if self.train_config.skip_first_sample:
|
if self.train_config.skip_first_sample:
|
||||||
self.print("Skipping first sample due to config setting")
|
self.print("Skipping first sample due to config setting")
|
||||||
else:
|
else:
|
||||||
self.print("Generating baseline samples before training")
|
self.print("Generating baseline samples before training")
|
||||||
self.sample(0, is_first=True)
|
self.sample(0)
|
||||||
|
|
||||||
self.progress_bar = tqdm(
|
self.progress_bar = tqdm(
|
||||||
total=self.train_config.steps,
|
total=self.train_config.steps,
|
||||||
desc=self.job.name,
|
desc=self.job.name,
|
||||||
leave=True
|
leave=True,
|
||||||
|
initial=self.step_num,
|
||||||
|
iterable=range(0, self.train_config.steps),
|
||||||
)
|
)
|
||||||
# set it to our current step in case it was updated from a load
|
|
||||||
self.progress_bar.update(self.step_num)
|
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
|
# self.step_num = 0
|
||||||
for step in range(self.step_num, self.train_config.steps):
|
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 ###
|
### HOOK ###
|
||||||
loss_dict = self.hook_train_loop()
|
loss_dict = self.hook_train_loop(batch)
|
||||||
|
flush()
|
||||||
|
|
||||||
if self.train_config.optimizer.lower().startswith('dadaptation'):
|
if self.train_config.optimizer.lower().startswith('dadaptation') or \
|
||||||
|
self.train_config.optimizer.lower().startswith('prodigy'):
|
||||||
learning_rate = (
|
learning_rate = (
|
||||||
optimizer.param_groups[0]["d"] *
|
optimizer.param_groups[0]["d"] *
|
||||||
optimizer.param_groups[0]["lr"]
|
optimizer.param_groups[0]["lr"]
|
||||||
|
|||||||
@@ -1,14 +1,24 @@
|
|||||||
|
from datetime import datetime
|
||||||
import os
|
import os
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import ForwardRef
|
from typing import TYPE_CHECKING, Union
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
from jobs.process.BaseProcess import BaseProcess
|
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):
|
class BaseTrainProcess(BaseProcess):
|
||||||
process_id: int
|
process_id: int
|
||||||
config: OrderedDict
|
config: OrderedDict
|
||||||
progress_bar: ForwardRef('tqdm') = None
|
writer: 'SummaryWriter'
|
||||||
|
job: Union['TrainJob', 'BaseJob', 'ExtensionJob']
|
||||||
|
progress_bar: 'tqdm' = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -18,11 +28,14 @@ class BaseTrainProcess(BaseProcess):
|
|||||||
):
|
):
|
||||||
super().__init__(process_id, job, config)
|
super().__init__(process_id, job, config)
|
||||||
self.progress_bar = None
|
self.progress_bar = None
|
||||||
self.writer = self.job.writer
|
self.writer = None
|
||||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
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.job.name)
|
self.save_root = os.path.join(self.training_folder, self.name)
|
||||||
self.step = 0
|
self.step = 0
|
||||||
self.first_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):
|
def run(self):
|
||||||
super().run()
|
super().run()
|
||||||
@@ -37,3 +50,19 @@ class BaseTrainProcess(BaseProcess):
|
|||||||
self.progress_bar.update()
|
self.progress_bar.update()
|
||||||
else:
|
else:
|
||||||
print(*args)
|
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()
|
||||||
101
jobs/process/ModRescaleLoraProcess.py
Normal file
101
jobs/process/ModRescaleLoraProcess.py
Normal file
@@ -0,0 +1,101 @@
|
|||||||
|
import gc
|
||||||
|
import os
|
||||||
|
from collections import OrderedDict
|
||||||
|
from typing import ForwardRef
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import save_file, load_file
|
||||||
|
|
||||||
|
from jobs.process.BaseProcess import BaseProcess
|
||||||
|
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.train_tools import get_torch_dtype
|
||||||
|
|
||||||
|
|
||||||
|
class ModRescaleLoraProcess(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)
|
||||||
|
self.input_path = self.get_conf('input_path', required=True)
|
||||||
|
self.output_path = self.get_conf('output_path', required=True)
|
||||||
|
self.replace_meta = self.get_conf('replace_meta', default=False)
|
||||||
|
self.save_dtype = self.get_conf('save_dtype', default='fp16', as_type=get_torch_dtype)
|
||||||
|
self.current_weight = self.get_conf('current_weight', required=True, as_type=float)
|
||||||
|
self.target_weight = self.get_conf('target_weight', required=True, as_type=float)
|
||||||
|
self.scale_target = self.get_conf('scale_target', default='up_down') # alpha or up_down
|
||||||
|
self.is_xl = self.get_conf('is_xl', default=False, as_type=bool)
|
||||||
|
self.is_v2 = self.get_conf('is_v2', default=False, as_type=bool)
|
||||||
|
|
||||||
|
self.progress_bar = None
|
||||||
|
|
||||||
|
def run(self):
|
||||||
|
super().run()
|
||||||
|
source_state_dict = load_file(self.input_path)
|
||||||
|
source_meta = load_metadata_from_safetensors(self.input_path)
|
||||||
|
|
||||||
|
if self.replace_meta:
|
||||||
|
self.meta.update(
|
||||||
|
add_base_model_info_to_meta(
|
||||||
|
self.meta,
|
||||||
|
is_xl=self.is_xl,
|
||||||
|
is_v2=self.is_v2,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
|
||||||
|
else:
|
||||||
|
save_meta = get_meta_for_safetensors(source_meta, self.job.name, add_software_info=False)
|
||||||
|
|
||||||
|
# save
|
||||||
|
os.makedirs(os.path.dirname(self.output_path), exist_ok=True)
|
||||||
|
|
||||||
|
new_state_dict = OrderedDict()
|
||||||
|
|
||||||
|
for key in list(source_state_dict.keys()):
|
||||||
|
v = source_state_dict[key]
|
||||||
|
v = v.detach().clone().to("cpu").to(get_torch_dtype('fp32'))
|
||||||
|
|
||||||
|
# all loras have an alpha, up weight and down weight
|
||||||
|
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.alpha",
|
||||||
|
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.lora_down.weight",
|
||||||
|
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.lora_up.weight",
|
||||||
|
# we can rescale by adjusting the alpha or the up weights, or the up and down weights
|
||||||
|
# I assume doing both up and down would be best all around, but I'm not sure
|
||||||
|
# some locons also have mid weights, we will leave those alone for now, will work without them
|
||||||
|
|
||||||
|
# when adjusting alpha, it is used to calculate the multiplier in a lora module
|
||||||
|
# - scale = alpha / lora_dim
|
||||||
|
# - output = layer_out + lora_up_out * multiplier * scale
|
||||||
|
total_module_scale = torch.tensor(self.current_weight / self.target_weight) \
|
||||||
|
.to("cpu", dtype=get_torch_dtype('fp32'))
|
||||||
|
num_modules_layers = 2 # up and down
|
||||||
|
up_down_scale = torch.pow(total_module_scale, 1.0 / num_modules_layers) \
|
||||||
|
.to("cpu", dtype=get_torch_dtype('fp32'))
|
||||||
|
# only update alpha
|
||||||
|
if self.scale_target == 'alpha' and key.endswith('.alpha'):
|
||||||
|
v = v * total_module_scale
|
||||||
|
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
|
||||||
|
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)
|
||||||
|
|
||||||
|
# cleanup incase there are other jobs
|
||||||
|
del new_state_dict
|
||||||
|
del source_state_dict
|
||||||
|
del source_meta
|
||||||
|
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
print(f"Saved to {self.output_path}")
|
||||||
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
|
return loss_dict
|
||||||
|
|
||||||
def hook_train_loop(self):
|
def hook_train_loop(self, batch):
|
||||||
if self.hack_config.type == 'suppression':
|
if self.hack_config.type == 'suppression':
|
||||||
return self.supress_loop()
|
return self.supress_loop()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,22 +1,14 @@
|
|||||||
# ref:
|
import glob
|
||||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
|
||||||
import time
|
|
||||||
from collections import OrderedDict
|
|
||||||
import os
|
import os
|
||||||
from typing import Optional
|
from collections import OrderedDict
|
||||||
|
import random
|
||||||
|
from typing import Optional, List
|
||||||
|
|
||||||
from safetensors.torch import load_file, save_file
|
from safetensors.torch import save_file, load_file
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from toolkit.config_modules import SliderConfig
|
|
||||||
from toolkit.layers import ReductionKernel
|
from toolkit.layers import ReductionKernel
|
||||||
from toolkit.paths import REPOS_ROOT
|
|
||||||
import sys
|
|
||||||
|
|
||||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
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, apply_noise_offset
|
||||||
import gc
|
import gc
|
||||||
from toolkit import train_tools
|
from toolkit import train_tools
|
||||||
@@ -38,12 +30,10 @@ class RescaleConfig:
|
|||||||
):
|
):
|
||||||
self.from_resolution = kwargs.get('from_resolution', 512)
|
self.from_resolution = kwargs.get('from_resolution', 512)
|
||||||
self.scale = kwargs.get('scale', 0.5)
|
self.scale = kwargs.get('scale', 0.5)
|
||||||
self.prompt_file = kwargs.get('prompt_file', None)
|
self.latent_tensor_dir = kwargs.get('latent_tensor_dir', None)
|
||||||
self.prompt_tensors = kwargs.get('prompt_tensors', None)
|
self.num_latent_tensors = kwargs.get('num_latent_tensors', 1000)
|
||||||
self.to_resolution = kwargs.get('to_resolution', int(self.from_resolution * self.scale))
|
self.to_resolution = kwargs.get('to_resolution', int(self.from_resolution * self.scale))
|
||||||
|
self.prompt_dropout = kwargs.get('prompt_dropout', 0.1)
|
||||||
if self.prompt_file is None:
|
|
||||||
raise ValueError("prompt_file is required")
|
|
||||||
|
|
||||||
|
|
||||||
class PromptEmbedsCache:
|
class PromptEmbedsCache:
|
||||||
@@ -61,12 +51,12 @@ class PromptEmbedsCache:
|
|||||||
|
|
||||||
class TrainSDRescaleProcess(BaseSDTrainProcess):
|
class TrainSDRescaleProcess(BaseSDTrainProcess):
|
||||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
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)
|
super().__init__(process_id, job, config)
|
||||||
self.step_num = 0
|
self.step_num = 0
|
||||||
self.start_step = 0
|
self.start_step = 0
|
||||||
self.device = self.get_conf('device', self.job.device)
|
self.device = self.get_conf('device', self.job.device)
|
||||||
self.device_torch = torch.device(self.device)
|
self.device_torch = torch.device(self.device)
|
||||||
self.prompt_cache = PromptEmbedsCache()
|
|
||||||
self.rescale_config = RescaleConfig(**self.get_conf('rescale', required=True))
|
self.rescale_config = RescaleConfig(**self.get_conf('rescale', required=True))
|
||||||
self.reduce_size_fn = ReductionKernel(
|
self.reduce_size_fn = ReductionKernel(
|
||||||
in_channels=4,
|
in_channels=4,
|
||||||
@@ -74,202 +64,211 @@ class TrainSDRescaleProcess(BaseSDTrainProcess):
|
|||||||
dtype=get_torch_dtype(self.train_config.dtype),
|
dtype=get_torch_dtype(self.train_config.dtype),
|
||||||
device=self.device_torch,
|
device=self.device_torch,
|
||||||
)
|
)
|
||||||
self.prompt_txt_list = []
|
|
||||||
|
self.latent_paths: List[str] = []
|
||||||
|
self.empty_embedding: PromptEmbeds = None
|
||||||
|
|
||||||
def before_model_load(self):
|
def before_model_load(self):
|
||||||
pass
|
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):
|
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
|
# Move train model encoder to cpu
|
||||||
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
|
|
||||||
if isinstance(self.sd.text_encoder, list):
|
if isinstance(self.sd.text_encoder, list):
|
||||||
for encoder in self.sd.text_encoder:
|
for encoder in self.sd.text_encoder:
|
||||||
encoder.to("cpu")
|
encoder.to('cpu')
|
||||||
|
encoder.eval()
|
||||||
|
encoder.requires_grad_(False)
|
||||||
else:
|
else:
|
||||||
self.sd.text_encoder.to("cpu")
|
self.sd.text_encoder.to('cpu')
|
||||||
self.prompt_cache = cache
|
self.sd.text_encoder.eval()
|
||||||
|
self.sd.text_encoder.requires_grad_(False)
|
||||||
|
|
||||||
|
# self.sd.unet.to('cpu')
|
||||||
|
flush()
|
||||||
|
|
||||||
|
self.get_latent_tensors()
|
||||||
|
|
||||||
flush()
|
flush()
|
||||||
# end hook_before_train_loop
|
# end hook_before_train_loop
|
||||||
|
|
||||||
def hook_train_loop(self):
|
def hook_train_loop(self, batch):
|
||||||
dtype = get_torch_dtype(self.train_config.dtype)
|
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)
|
|
||||||
neutral = self.prompt_cache[""].to(device=self.device_torch, dtype=dtype)
|
|
||||||
if prompt is None:
|
|
||||||
raise ValueError(f"Prompt {prompt_txt} is not in cache")
|
|
||||||
|
|
||||||
prompt_batch = train_tools.concat_prompt_embeddings(
|
|
||||||
prompt,
|
|
||||||
neutral,
|
|
||||||
self.train_config.batch_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
noise_scheduler = self.sd.noise_scheduler
|
|
||||||
optimizer = self.optimizer
|
|
||||||
lr_scheduler = self.lr_scheduler
|
|
||||||
loss_function = torch.nn.MSELoss()
|
loss_function = torch.nn.MSELoss()
|
||||||
|
|
||||||
def get_noise_pred(p, n, gs, cts, dn):
|
# train it
|
||||||
return self.predict_noise(
|
# Begin gradient accumulation
|
||||||
latents=dn,
|
self.sd.unet.train()
|
||||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
self.sd.unet.requires_grad_(True)
|
||||||
p, # unconditional
|
self.sd.unet.to(self.device_torch, dtype=dtype)
|
||||||
n, # positive
|
|
||||||
self.train_config.batch_size,
|
|
||||||
),
|
|
||||||
timestep=cts,
|
|
||||||
guidance_scale=gs,
|
|
||||||
)
|
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
self.sd.noise_scheduler.set_timesteps(
|
|
||||||
self.train_config.max_denoising_steps, device=self.device_torch
|
|
||||||
)
|
|
||||||
|
|
||||||
self.optimizer.zero_grad()
|
self.optimizer.zero_grad()
|
||||||
|
|
||||||
# # ger a random number of steps
|
# pick random latent tensor
|
||||||
timesteps_to = torch.randint(
|
latent_path = random.choice(self.latent_paths)
|
||||||
1, self.train_config.max_denoising_steps, (1,)
|
latent_tensor = load_file(latent_path)
|
||||||
).item()
|
|
||||||
|
|
||||||
# get noise
|
noise_pred_target = (latent_tensor['noise_pred_target']).to(self.device_torch, dtype=dtype)
|
||||||
noise = self.get_latent_noise(
|
latents = (latent_tensor['latents']).to(self.device_torch, dtype=dtype)
|
||||||
pixel_height=self.rescale_config.from_resolution,
|
guidance_scale = (latent_tensor['guidance_scale']).item()
|
||||||
pixel_width=self.rescale_config.from_resolution,
|
timestep = int((latent_tensor['timestep']).item())
|
||||||
).to(self.device_torch, dtype=dtype)
|
timesteps_to = int((latent_tensor['timesteps_to']).item())
|
||||||
|
# seed = int((latent_tensor['seed']).item())
|
||||||
|
|
||||||
# get latents
|
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||||
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
self.empty_embedding, # unconditional (negative prompt)
|
||||||
latents = latents.to(self.device_torch, dtype=dtype)
|
self.empty_embedding, # conditional (positive prompt)
|
||||||
#
|
self.train_config.batch_size,
|
||||||
# # predict without network
|
)
|
||||||
# assert self.network.is_active is False
|
self.sd.noise_scheduler.set_timesteps(
|
||||||
# denoised_latents = self.diffuse_some_steps(
|
timesteps_to, device=self.device_torch
|
||||||
# latents, # pass simple noise latents
|
|
||||||
# prompt_batch,
|
|
||||||
# start_timesteps=0,
|
|
||||||
# total_timesteps=timesteps_to,
|
|
||||||
# guidance_scale=3,
|
|
||||||
# )
|
|
||||||
# noise_scheduler.set_timesteps(1000)
|
|
||||||
#
|
|
||||||
# current_timestep = noise_scheduler.timesteps[
|
|
||||||
# int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
|
||||||
# ]
|
|
||||||
|
|
||||||
current_timestep = 0
|
|
||||||
denoised_latents = latents
|
|
||||||
# get noise prediction at full scale
|
|
||||||
from_prediction = get_noise_pred(
|
|
||||||
prompt, neutral, 1, current_timestep, denoised_latents
|
|
||||||
)
|
)
|
||||||
|
|
||||||
reduced_from_prediction = self.reduce_size_fn(from_prediction).to("cpu", dtype=torch.float32)
|
denoised_target = self.sd.noise_scheduler.step(noise_pred_target, timestep, latents).prev_sample
|
||||||
|
|
||||||
# get noise prediction at reduced scale
|
# get the reduced latents
|
||||||
to_denoised_latents = self.reduce_size_fn(denoised_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())
|
||||||
|
|
||||||
# start gradient
|
denoised_target.requires_grad = False
|
||||||
optimizer.zero_grad()
|
self.optimizer.zero_grad()
|
||||||
self.network.multiplier = 1.0
|
noise_pred_train = self.sd.predict_noise(
|
||||||
with self.network:
|
reduced_latents,
|
||||||
assert self.network.is_active is True
|
text_embeddings=text_embeddings,
|
||||||
to_prediction = get_noise_pred(
|
timestep=timestep,
|
||||||
prompt, neutral, 1, current_timestep, to_denoised_latents
|
guidance_scale=guidance_scale
|
||||||
).to("cpu", dtype=torch.float32)
|
|
||||||
|
|
||||||
reduced_from_prediction.requires_grad = False
|
|
||||||
from_prediction.requires_grad = False
|
|
||||||
|
|
||||||
loss = loss_function(
|
|
||||||
reduced_from_prediction,
|
|
||||||
to_prediction,
|
|
||||||
)
|
)
|
||||||
|
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_float = loss.item()
|
||||||
|
|
||||||
loss = loss.to(self.device_torch)
|
|
||||||
|
|
||||||
loss.backward()
|
loss.backward()
|
||||||
optimizer.step()
|
self.optimizer.step()
|
||||||
lr_scheduler.step()
|
self.lr_scheduler.step()
|
||||||
|
self.optimizer.zero_grad()
|
||||||
|
|
||||||
del (
|
|
||||||
reduced_from_prediction,
|
|
||||||
from_prediction,
|
|
||||||
to_denoised_latents,
|
|
||||||
to_prediction,
|
|
||||||
latents,
|
|
||||||
)
|
|
||||||
flush()
|
flush()
|
||||||
|
|
||||||
# reset network
|
|
||||||
self.network.multiplier = 1.0
|
|
||||||
|
|
||||||
loss_dict = OrderedDict(
|
loss_dict = OrderedDict(
|
||||||
{'loss': loss_float},
|
{'loss': loss_float},
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,30 +1,31 @@
|
|||||||
# ref:
|
# ref:
|
||||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||||
import time
|
import random
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
import os
|
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.config_modules import SliderConfig
|
||||||
|
from toolkit.layers import CheckpointGradients
|
||||||
from toolkit.paths import REPOS_ROOT
|
from toolkit.paths import REPOS_ROOT
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||||
|
from toolkit.train_tools import get_torch_dtype
|
||||||
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
|
import gc
|
||||||
from toolkit import train_tools
|
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
|
import torch
|
||||||
from leco import train_util, model_util
|
from .BaseSDTrainProcess import BaseSDTrainProcess
|
||||||
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
|
|
||||||
|
|
||||||
|
|
||||||
class ACTION_TYPES_SLIDER:
|
|
||||||
ERASE_NEGATIVE = 0
|
|
||||||
ENHANCE_NEGATIVE = 1
|
|
||||||
|
|
||||||
|
|
||||||
def flush():
|
def flush():
|
||||||
@@ -32,58 +33,10 @@ def flush():
|
|||||||
gc.collect()
|
gc.collect()
|
||||||
|
|
||||||
|
|
||||||
class EncodedPromptPair:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
target_class,
|
|
||||||
positive,
|
|
||||||
negative,
|
|
||||||
neutral,
|
|
||||||
width=512,
|
|
||||||
height=512,
|
|
||||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
|
||||||
multiplier=1.0,
|
|
||||||
weight=1.0
|
|
||||||
):
|
|
||||||
self.target_class = target_class
|
|
||||||
self.positive = positive
|
|
||||||
self.negative = negative
|
|
||||||
self.neutral = neutral
|
|
||||||
self.width = width
|
|
||||||
self.height = height
|
|
||||||
self.action: int = action
|
|
||||||
self.multiplier = multiplier
|
|
||||||
self.weight = weight
|
|
||||||
|
|
||||||
|
|
||||||
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):
|
class TrainSliderProcess(BaseSDTrainProcess):
|
||||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||||
super().__init__(process_id, job, config)
|
super().__init__(process_id, job, config)
|
||||||
|
self.prompt_txt_list = None
|
||||||
self.step_num = 0
|
self.step_num = 0
|
||||||
self.start_step = 0
|
self.start_step = 0
|
||||||
self.device = self.get_conf('device', self.job.device)
|
self.device = self.get_conf('device', self.job.device)
|
||||||
@@ -92,102 +45,87 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
self.prompt_cache = PromptEmbedsCache()
|
self.prompt_cache = PromptEmbedsCache()
|
||||||
self.prompt_pairs: list[EncodedPromptPair] = []
|
self.prompt_pairs: list[EncodedPromptPair] = []
|
||||||
self.anchor_pairs: list[EncodedAnchor] = []
|
self.anchor_pairs: list[EncodedAnchor] = []
|
||||||
|
# keep track of prompt chunk size
|
||||||
|
self.prompt_chunk_size = 1
|
||||||
|
|
||||||
def before_model_load(self):
|
def before_model_load(self):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def hook_before_train_loop(self):
|
def hook_before_train_loop(self):
|
||||||
|
|
||||||
|
# read line by line from file
|
||||||
|
if self.slider_config.prompt_file:
|
||||||
|
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"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
|
||||||
|
self.prompt_txt_list = self.prompt_txt_list[:self.train_config.steps]
|
||||||
|
# trim list to our max steps
|
||||||
|
|
||||||
cache = PromptEmbedsCache()
|
cache = PromptEmbedsCache()
|
||||||
prompt_pairs: list[EncodedPromptPair] = []
|
|
||||||
|
|
||||||
# get encoded latents for our prompts
|
# get encoded latents for our prompts
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
neutral = ""
|
# list of neutrals. Can come from file or be empty
|
||||||
for target in self.slider_config.targets:
|
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
|
||||||
# build the cache
|
|
||||||
for prompt in [
|
|
||||||
target.target_class,
|
|
||||||
target.positive,
|
|
||||||
target.negative,
|
|
||||||
neutral # empty neutral
|
|
||||||
]:
|
|
||||||
if cache[prompt] is None:
|
|
||||||
cache[prompt] = self.sd.encode_prompt(prompt)
|
|
||||||
for resolution in self.slider_config.resolutions:
|
|
||||||
width, height = resolution
|
|
||||||
erase_negative = len(target.positive.strip()) == 0
|
|
||||||
enhance_positive = len(target.negative.strip()) == 0
|
|
||||||
|
|
||||||
both = not erase_negative and not enhance_positive
|
# 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 erase_negative and enhance_positive:
|
# remove duplicates
|
||||||
raise ValueError("target must have at least one of positive or negative or both")
|
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
|
||||||
# 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 both or erase_negative:
|
# encode them
|
||||||
prompt_pairs += [
|
cache = encode_prompts_to_cache(
|
||||||
# erase standard
|
prompt_list=prompts_to_cache,
|
||||||
EncodedPromptPair(
|
sd=self.sd,
|
||||||
target_class=cache[target.target_class],
|
cache=cache,
|
||||||
positive=cache[target.positive],
|
prompt_tensor_file=self.slider_config.prompt_tensors
|
||||||
negative=cache[target.negative],
|
)
|
||||||
neutral=cache[neutral],
|
|
||||||
width=width,
|
prompt_pairs = []
|
||||||
height=height,
|
prompt_batches = []
|
||||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
|
||||||
multiplier=target.multiplier,
|
for target in self.slider_config.targets:
|
||||||
weight=target.weight
|
prompt_pair_batch = build_prompt_pair_batch_from_cache(
|
||||||
),
|
cache=cache,
|
||||||
]
|
target=target,
|
||||||
if both or enhance_positive:
|
neutral=neutral,
|
||||||
prompt_pairs += [
|
|
||||||
# enhance standard, swap pos neg
|
)
|
||||||
EncodedPromptPair(
|
if self.slider_config.batch_full_slide:
|
||||||
target_class=cache[target.target_class],
|
# concat the prompt pairs
|
||||||
positive=cache[target.negative],
|
# this allows us to run the entire 4 part process in one shot (for slider)
|
||||||
negative=cache[target.positive],
|
self.prompt_chunk_size = 4
|
||||||
neutral=cache[neutral],
|
concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
|
||||||
width=width,
|
prompt_pairs += [concat_prompt_pair_batch]
|
||||||
height=height,
|
else:
|
||||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
self.prompt_chunk_size = 1
|
||||||
multiplier=target.multiplier,
|
# do them one at a time (probably not necessary after new optimizations)
|
||||||
weight=target.weight
|
prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
|
||||||
),
|
|
||||||
]
|
|
||||||
if both or enhance_positive:
|
|
||||||
prompt_pairs += [
|
|
||||||
# erase inverted
|
|
||||||
EncodedPromptPair(
|
|
||||||
target_class=cache[target.target_class],
|
|
||||||
positive=cache[target.negative],
|
|
||||||
negative=cache[target.positive],
|
|
||||||
neutral=cache[neutral],
|
|
||||||
width=width,
|
|
||||||
height=height,
|
|
||||||
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
|
||||||
multiplier=target.multiplier * -1.0,
|
|
||||||
weight=target.weight
|
|
||||||
),
|
|
||||||
]
|
|
||||||
if both or erase_negative:
|
|
||||||
prompt_pairs += [
|
|
||||||
# enhance inverted
|
|
||||||
EncodedPromptPair(
|
|
||||||
target_class=cache[target.target_class],
|
|
||||||
positive=cache[target.positive],
|
|
||||||
negative=cache[target.negative],
|
|
||||||
neutral=cache[neutral],
|
|
||||||
width=width,
|
|
||||||
height=height,
|
|
||||||
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
|
||||||
multiplier=target.multiplier * -1.0,
|
|
||||||
weight=target.weight
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
# setup anchors
|
# setup anchors
|
||||||
anchor_pairs = []
|
anchor_pairs = []
|
||||||
@@ -200,13 +138,26 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
if cache[prompt] == None:
|
if cache[prompt] == None:
|
||||||
cache[prompt] = self.sd.encode_prompt(prompt)
|
cache[prompt] = self.sd.encode_prompt(prompt)
|
||||||
|
|
||||||
|
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 += [
|
anchor_pairs += [
|
||||||
EncodedAnchor(
|
concat_anchors(anchor_batch).to('cpu')
|
||||||
prompt=cache[anchor.prompt],
|
|
||||||
neg_prompt=cache[anchor.neg_prompt],
|
|
||||||
multiplier=anchor.multiplier
|
|
||||||
)
|
|
||||||
]
|
]
|
||||||
|
if len(anchor_pairs) > 0:
|
||||||
|
self.anchor_pairs = anchor_pairs
|
||||||
|
|
||||||
# move to cpu to save vram
|
# move to cpu to save vram
|
||||||
# We don't need text encoder anymore, but keep it on cpu for sampling
|
# We don't need text encoder anymore, but keep it on cpu for sampling
|
||||||
@@ -218,48 +169,45 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
self.sd.text_encoder.to("cpu")
|
self.sd.text_encoder.to("cpu")
|
||||||
self.prompt_cache = cache
|
self.prompt_cache = cache
|
||||||
self.prompt_pairs = prompt_pairs
|
self.prompt_pairs = prompt_pairs
|
||||||
self.anchor_pairs = anchor_pairs
|
# self.anchor_pairs = anchor_pairs
|
||||||
flush()
|
flush()
|
||||||
# end hook_before_train_loop
|
# end hook_before_train_loop
|
||||||
|
|
||||||
def hook_train_loop(self):
|
def hook_train_loop(self, batch):
|
||||||
dtype = get_torch_dtype(self.train_config.dtype)
|
dtype = get_torch_dtype(self.train_config.dtype)
|
||||||
|
|
||||||
# get a random pair
|
# get a random pair
|
||||||
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
||||||
torch.randint(0, len(self.prompt_pairs), (1,)).item()
|
torch.randint(0, len(self.prompt_pairs), (1,)).item()
|
||||||
]
|
]
|
||||||
|
# move to device and dtype
|
||||||
|
prompt_pair.to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
height = prompt_pair.height
|
# get a random resolution
|
||||||
width = prompt_pair.width
|
height, width = self.slider_config.resolutions[
|
||||||
target_class = prompt_pair.target_class
|
torch.randint(0, len(self.slider_config.resolutions), (1,)).item()
|
||||||
neutral = prompt_pair.neutral
|
]
|
||||||
negative = prompt_pair.negative
|
if self.train_config.gradient_checkpointing:
|
||||||
positive = prompt_pair.positive
|
# may get disabled elsewhere
|
||||||
weight = prompt_pair.weight
|
self.sd.unet.enable_gradient_checkpointing()
|
||||||
multiplier = prompt_pair.multiplier
|
|
||||||
|
|
||||||
unet = self.sd.unet
|
|
||||||
noise_scheduler = self.sd.noise_scheduler
|
noise_scheduler = self.sd.noise_scheduler
|
||||||
optimizer = self.optimizer
|
optimizer = self.optimizer
|
||||||
lr_scheduler = self.lr_scheduler
|
lr_scheduler = self.lr_scheduler
|
||||||
loss_function = torch.nn.MSELoss()
|
loss_function = torch.nn.MSELoss()
|
||||||
|
|
||||||
def get_noise_pred(p, n, gs, cts, dn):
|
def get_noise_pred(neg, pos, gs, cts, dn):
|
||||||
return self.predict_noise(
|
return self.sd.predict_noise(
|
||||||
latents=dn,
|
latents=dn,
|
||||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||||
p, # unconditional
|
neg, # negative prompt
|
||||||
n, # positive
|
pos, # positive prompt
|
||||||
self.train_config.batch_size,
|
self.train_config.batch_size,
|
||||||
),
|
),
|
||||||
timestep=cts,
|
timestep=cts,
|
||||||
guidance_scale=gs,
|
guidance_scale=gs,
|
||||||
)
|
)
|
||||||
|
|
||||||
# set network multiplier
|
|
||||||
self.network.multiplier = multiplier
|
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
self.sd.noise_scheduler.set_timesteps(
|
self.sd.noise_scheduler.set_timesteps(
|
||||||
self.train_config.max_denoising_steps, device=self.device_torch
|
self.train_config.max_denoising_steps, device=self.device_torch
|
||||||
@@ -272,10 +220,15 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
1, self.train_config.max_denoising_steps, (1,)
|
1, self.train_config.max_denoising_steps, (1,)
|
||||||
).item()
|
).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
|
# get noise
|
||||||
noise = self.get_latent_noise(
|
noise = self.sd.get_latent_noise(
|
||||||
pixel_height=height,
|
pixel_height=height,
|
||||||
pixel_width=width,
|
pixel_width=width,
|
||||||
|
batch_size=true_batch_size,
|
||||||
|
noise_offset=self.train_config.noise_offset,
|
||||||
).to(self.device_torch, dtype=dtype)
|
).to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
# get latents
|
# get latents
|
||||||
@@ -284,12 +237,13 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
|
|
||||||
with self.network:
|
with self.network:
|
||||||
assert self.network.is_active
|
assert self.network.is_active
|
||||||
self.network.multiplier = multiplier
|
# pass the multiplier list to the network
|
||||||
denoised_latents = self.diffuse_some_steps(
|
self.network.multiplier = prompt_pair.multiplier_list
|
||||||
|
denoised_latents = self.sd.diffuse_some_steps(
|
||||||
latents, # pass simple noise latents
|
latents, # pass simple noise latents
|
||||||
train_tools.concat_prompt_embeddings(
|
train_tools.concat_prompt_embeddings(
|
||||||
positive, # unconditional
|
prompt_pair.positive_target, # unconditional
|
||||||
target_class, # target
|
prompt_pair.target_class, # target
|
||||||
self.train_config.batch_size,
|
self.train_config.batch_size,
|
||||||
),
|
),
|
||||||
start_timesteps=0,
|
start_timesteps=0,
|
||||||
@@ -297,100 +251,183 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
guidance_scale=3,
|
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)
|
noise_scheduler.set_timesteps(1000)
|
||||||
|
|
||||||
current_timestep = noise_scheduler.timesteps[
|
current_timestep = noise_scheduler.timesteps[
|
||||||
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
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(
|
positive_latents = get_noise_pred(
|
||||||
positive, negative, 1, current_timestep, denoised_latents
|
prompt_pair.positive_target, # negative prompt
|
||||||
).to("cpu", dtype=torch.float32)
|
prompt_pair.negative_target, # positive prompt
|
||||||
|
1,
|
||||||
|
current_timestep,
|
||||||
|
denoised_latents
|
||||||
|
)
|
||||||
|
positive_latents.requires_grad = False
|
||||||
|
positive_latents_chunks = torch.chunk(positive_latents, self.prompt_chunk_size, dim=0)
|
||||||
|
|
||||||
neutral_latents = get_noise_pred(
|
neutral_latents = get_noise_pred(
|
||||||
positive, neutral, 1, current_timestep, denoised_latents
|
prompt_pair.positive_target, # negative prompt
|
||||||
).to("cpu", dtype=torch.float32)
|
prompt_pair.empty_prompt, # positive prompt (normally neutral
|
||||||
|
1,
|
||||||
|
current_timestep,
|
||||||
|
denoised_latents
|
||||||
|
)
|
||||||
|
neutral_latents.requires_grad = False
|
||||||
|
neutral_latents_chunks = torch.chunk(neutral_latents, self.prompt_chunk_size, dim=0)
|
||||||
|
|
||||||
unconditional_latents = get_noise_pred(
|
unconditional_latents = get_noise_pred(
|
||||||
positive, positive, 1, current_timestep, denoised_latents
|
prompt_pair.positive_target, # negative prompt
|
||||||
).to("cpu", dtype=torch.float32)
|
prompt_pair.positive_target, # positive prompt
|
||||||
|
1,
|
||||||
anchor_loss = None
|
current_timestep,
|
||||||
if len(self.anchor_pairs) > 0:
|
denoised_latents
|
||||||
# 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
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
with self.network:
|
|
||||||
self.network.multiplier = prompt_pair.multiplier
|
|
||||||
target_latents = get_noise_pred(
|
|
||||||
positive, 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
|
unconditional_latents.requires_grad = False
|
||||||
guidance_scale = 1.0
|
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
|
# 4.20 GB RAM for 512x512
|
||||||
if erase:
|
anchor_loss_float = None
|
||||||
offset_neutral -= offset
|
if len(self.anchor_pairs) > 0:
|
||||||
else:
|
with torch.no_grad():
|
||||||
# enhance
|
# get a random anchor pair
|
||||||
offset_neutral += offset
|
anchor: EncodedAnchor = self.anchor_pairs[
|
||||||
|
torch.randint(0, len(self.anchor_pairs), (1,)).item()
|
||||||
|
]
|
||||||
|
anchor.to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
loss = loss_function(
|
# first we get the target prediction without network active
|
||||||
target_latents,
|
anchor_target_noise = get_noise_pred(
|
||||||
offset_neutral,
|
anchor.neg_prompt, anchor.prompt, 1, current_timestep, denoised_latents
|
||||||
) * weight
|
# ).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:
|
# 4.32 GB RAM for 512x512
|
||||||
loss += anchor_loss
|
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()
|
optimizer.step()
|
||||||
lr_scheduler.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 (
|
del (
|
||||||
positive_latents,
|
positive_latents,
|
||||||
neutral_latents,
|
neutral_latents,
|
||||||
unconditional_latents,
|
unconditional_latents,
|
||||||
target_latents,
|
latents
|
||||||
latents,
|
|
||||||
)
|
)
|
||||||
|
# move back to cpu
|
||||||
|
prompt_pair.to("cpu")
|
||||||
flush()
|
flush()
|
||||||
|
|
||||||
# reset network
|
# reset network
|
||||||
@@ -399,9 +436,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
|||||||
loss_dict = OrderedDict(
|
loss_dict = OrderedDict(
|
||||||
{'loss': loss_float},
|
{'loss': loss_float},
|
||||||
)
|
)
|
||||||
if anchor_loss is not None:
|
if anchor_loss_float is not None:
|
||||||
loss_dict['sl_l'] = loss_slide
|
loss_dict['sl_l'] = loss_float
|
||||||
loss_dict['an_l'] = anchor_loss.item()
|
loss_dict['an_l'] = anchor_loss_float
|
||||||
|
|
||||||
return loss_dict
|
return loss_dict
|
||||||
# end hook_train_loop
|
# end hook_train_loop
|
||||||
|
|||||||
408
jobs/process/TrainSliderProcessOld.py
Normal file
408
jobs/process/TrainSliderProcessOld.py
Normal file
@@ -0,0 +1,408 @@
|
|||||||
|
# ref:
|
||||||
|
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||||
|
import time
|
||||||
|
from collections import OrderedDict
|
||||||
|
import os
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from toolkit.config_modules import SliderConfig
|
||||||
|
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
|
||||||
|
import gc
|
||||||
|
from toolkit import train_tools
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from leco import train_util, model_util
|
||||||
|
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
|
||||||
|
|
||||||
|
|
||||||
|
class ACTION_TYPES_SLIDER:
|
||||||
|
ERASE_NEGATIVE = 0
|
||||||
|
ENHANCE_NEGATIVE = 1
|
||||||
|
|
||||||
|
|
||||||
|
def flush():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
|
||||||
|
class EncodedPromptPair:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
target_class,
|
||||||
|
positive,
|
||||||
|
negative,
|
||||||
|
neutral,
|
||||||
|
width=512,
|
||||||
|
height=512,
|
||||||
|
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||||
|
multiplier=1.0,
|
||||||
|
weight=1.0
|
||||||
|
):
|
||||||
|
self.target_class = target_class
|
||||||
|
self.positive = positive
|
||||||
|
self.negative = negative
|
||||||
|
self.neutral = neutral
|
||||||
|
self.width = width
|
||||||
|
self.height = height
|
||||||
|
self.action: int = action
|
||||||
|
self.multiplier = multiplier
|
||||||
|
self.weight = weight
|
||||||
|
|
||||||
|
|
||||||
|
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 TrainSliderProcessOld(BaseSDTrainProcess):
|
||||||
|
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||||
|
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.slider_config = SliderConfig(**self.get_conf('slider', {}))
|
||||||
|
self.prompt_cache = PromptEmbedsCache()
|
||||||
|
self.prompt_pairs: list[EncodedPromptPair] = []
|
||||||
|
self.anchor_pairs: list[EncodedAnchor] = []
|
||||||
|
|
||||||
|
def before_model_load(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def hook_before_train_loop(self):
|
||||||
|
cache = PromptEmbedsCache()
|
||||||
|
prompt_pairs: list[EncodedPromptPair] = []
|
||||||
|
|
||||||
|
# get encoded latents for our prompts
|
||||||
|
with torch.no_grad():
|
||||||
|
neutral = ""
|
||||||
|
for target in self.slider_config.targets:
|
||||||
|
# build the cache
|
||||||
|
for prompt in [
|
||||||
|
target.target_class,
|
||||||
|
target.positive,
|
||||||
|
target.negative,
|
||||||
|
neutral # empty neutral
|
||||||
|
]:
|
||||||
|
if cache[prompt] is None:
|
||||||
|
cache[prompt] = self.sd.encode_prompt(prompt)
|
||||||
|
for resolution in self.slider_config.resolutions:
|
||||||
|
width, height = resolution
|
||||||
|
only_erase = len(target.positive.strip()) == 0
|
||||||
|
only_enhance = len(target.negative.strip()) == 0
|
||||||
|
|
||||||
|
both = not only_erase and not only_enhance
|
||||||
|
|
||||||
|
if only_erase and only_enhance:
|
||||||
|
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 both or only_erase:
|
||||||
|
prompt_pairs += [
|
||||||
|
# erase standard
|
||||||
|
EncodedPromptPair(
|
||||||
|
target_class=cache[target.target_class],
|
||||||
|
positive=cache[target.positive],
|
||||||
|
negative=cache[target.negative],
|
||||||
|
neutral=cache[neutral],
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||||
|
multiplier=target.multiplier,
|
||||||
|
weight=target.weight
|
||||||
|
),
|
||||||
|
]
|
||||||
|
if both or only_enhance:
|
||||||
|
prompt_pairs += [
|
||||||
|
# enhance standard, swap pos neg
|
||||||
|
EncodedPromptPair(
|
||||||
|
target_class=cache[target.target_class],
|
||||||
|
positive=cache[target.negative],
|
||||||
|
negative=cache[target.positive],
|
||||||
|
neutral=cache[neutral],
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||||
|
multiplier=target.multiplier,
|
||||||
|
weight=target.weight
|
||||||
|
),
|
||||||
|
]
|
||||||
|
if both:
|
||||||
|
prompt_pairs += [
|
||||||
|
# erase inverted
|
||||||
|
EncodedPromptPair(
|
||||||
|
target_class=cache[target.target_class],
|
||||||
|
positive=cache[target.negative],
|
||||||
|
negative=cache[target.positive],
|
||||||
|
neutral=cache[neutral],
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
|
||||||
|
multiplier=target.multiplier * -1.0,
|
||||||
|
weight=target.weight
|
||||||
|
),
|
||||||
|
]
|
||||||
|
prompt_pairs += [
|
||||||
|
# enhance inverted
|
||||||
|
EncodedPromptPair(
|
||||||
|
target_class=cache[target.target_class],
|
||||||
|
positive=cache[target.positive],
|
||||||
|
negative=cache[target.negative],
|
||||||
|
neutral=cache[neutral],
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
|
||||||
|
multiplier=target.multiplier * -1.0,
|
||||||
|
weight=target.weight
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
# setup anchors
|
||||||
|
anchor_pairs = []
|
||||||
|
for anchor in self.slider_config.anchors:
|
||||||
|
# build the cache
|
||||||
|
for prompt in [
|
||||||
|
anchor.prompt,
|
||||||
|
anchor.neg_prompt # empty neutral
|
||||||
|
]:
|
||||||
|
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
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
# 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
|
||||||
|
if isinstance(self.sd.text_encoder, list):
|
||||||
|
for encoder in self.sd.text_encoder:
|
||||||
|
encoder.to("cpu")
|
||||||
|
else:
|
||||||
|
self.sd.text_encoder.to("cpu")
|
||||||
|
self.prompt_cache = cache
|
||||||
|
self.prompt_pairs = prompt_pairs
|
||||||
|
self.anchor_pairs = anchor_pairs
|
||||||
|
flush()
|
||||||
|
# end hook_before_train_loop
|
||||||
|
|
||||||
|
def hook_train_loop(self, batch):
|
||||||
|
dtype = get_torch_dtype(self.train_config.dtype)
|
||||||
|
|
||||||
|
# get a random pair
|
||||||
|
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
||||||
|
torch.randint(0, len(self.prompt_pairs), (1,)).item()
|
||||||
|
]
|
||||||
|
|
||||||
|
height = prompt_pair.height
|
||||||
|
width = prompt_pair.width
|
||||||
|
target_class = prompt_pair.target_class
|
||||||
|
neutral = prompt_pair.neutral
|
||||||
|
negative = prompt_pair.negative
|
||||||
|
positive = prompt_pair.positive
|
||||||
|
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(p, n, gs, cts, dn):
|
||||||
|
return self.sd.predict_noise(
|
||||||
|
latents=dn,
|
||||||
|
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||||
|
p, # unconditional
|
||||||
|
n, # positive
|
||||||
|
self.train_config.batch_size,
|
||||||
|
),
|
||||||
|
timestep=cts,
|
||||||
|
guidance_scale=gs,
|
||||||
|
)
|
||||||
|
|
||||||
|
# set network multiplier
|
||||||
|
self.network.multiplier = multiplier
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
self.sd.noise_scheduler.set_timesteps(
|
||||||
|
self.train_config.max_denoising_steps, device=self.device_torch
|
||||||
|
)
|
||||||
|
|
||||||
|
self.optimizer.zero_grad()
|
||||||
|
|
||||||
|
# ger a random number of steps
|
||||||
|
timesteps_to = torch.randint(
|
||||||
|
1, self.train_config.max_denoising_steps, (1,)
|
||||||
|
).item()
|
||||||
|
|
||||||
|
# get 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
|
||||||
|
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
||||||
|
latents = latents.to(self.device_torch, dtype=dtype)
|
||||||
|
|
||||||
|
with self.network:
|
||||||
|
assert self.network.is_active
|
||||||
|
self.network.multiplier = multiplier
|
||||||
|
denoised_latents = self.sd.diffuse_some_steps(
|
||||||
|
latents, # pass simple noise latents
|
||||||
|
train_tools.concat_prompt_embeddings(
|
||||||
|
positive, # unconditional
|
||||||
|
target_class, # target
|
||||||
|
self.train_config.batch_size,
|
||||||
|
),
|
||||||
|
start_timesteps=0,
|
||||||
|
total_timesteps=timesteps_to,
|
||||||
|
guidance_scale=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
noise_scheduler.set_timesteps(1000)
|
||||||
|
|
||||||
|
current_timestep = noise_scheduler.timesteps[
|
||||||
|
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
||||||
|
]
|
||||||
|
|
||||||
|
positive_latents = get_noise_pred(
|
||||||
|
positive, negative, 1, current_timestep, denoised_latents
|
||||||
|
).to("cpu", dtype=torch.float32)
|
||||||
|
|
||||||
|
neutral_latents = get_noise_pred(
|
||||||
|
positive, neutral, 1, current_timestep, denoised_latents
|
||||||
|
).to("cpu", dtype=torch.float32)
|
||||||
|
|
||||||
|
unconditional_latents = get_noise_pred(
|
||||||
|
positive, positive, 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
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
with self.network:
|
||||||
|
self.network.multiplier = prompt_pair.multiplier
|
||||||
|
target_latents = get_noise_pred(
|
||||||
|
positive, 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
|
||||||
|
|
||||||
|
offset = guidance_scale * (positive_latents - unconditional_latents)
|
||||||
|
|
||||||
|
offset_neutral = neutral_latents
|
||||||
|
if erase:
|
||||||
|
offset_neutral -= offset
|
||||||
|
else:
|
||||||
|
# enhance
|
||||||
|
offset_neutral += offset
|
||||||
|
|
||||||
|
loss = loss_function(
|
||||||
|
target_latents,
|
||||||
|
offset_neutral,
|
||||||
|
) * weight
|
||||||
|
|
||||||
|
loss_slide = loss.item()
|
||||||
|
|
||||||
|
if anchor_loss is not None:
|
||||||
|
loss += anchor_loss
|
||||||
|
|
||||||
|
loss_float = loss.item()
|
||||||
|
|
||||||
|
loss = loss.to(self.device_torch)
|
||||||
|
|
||||||
|
loss.backward()
|
||||||
|
optimizer.step()
|
||||||
|
lr_scheduler.step()
|
||||||
|
|
||||||
|
del (
|
||||||
|
positive_latents,
|
||||||
|
neutral_latents,
|
||||||
|
unconditional_latents,
|
||||||
|
target_latents,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
flush()
|
||||||
|
|
||||||
|
# reset network
|
||||||
|
self.network.multiplier = 1.0
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
return loss_dict
|
||||||
|
# end hook_train_loop
|
||||||
@@ -24,6 +24,7 @@ from diffusers import AutoencoderKL
|
|||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
import time
|
import time
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
from .models.vgg19_critic import Critic
|
||||||
|
|
||||||
IMAGE_TRANSFORMS = transforms.Compose(
|
IMAGE_TRANSFORMS = transforms.Compose(
|
||||||
[
|
[
|
||||||
@@ -37,145 +38,6 @@ def unnormalize(tensor):
|
|||||||
return (tensor / 2 + 0.5).clamp(0, 1)
|
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):
|
class TrainVAEProcess(BaseTrainProcess):
|
||||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||||
super().__init__(process_id, job, config)
|
super().__init__(process_id, job, config)
|
||||||
|
|||||||
@@ -6,5 +6,11 @@ from .BaseTrainProcess import BaseTrainProcess
|
|||||||
from .TrainVAEProcess import TrainVAEProcess
|
from .TrainVAEProcess import TrainVAEProcess
|
||||||
from .BaseMergeProcess import BaseMergeProcess
|
from .BaseMergeProcess import BaseMergeProcess
|
||||||
from .TrainSliderProcess import TrainSliderProcess
|
from .TrainSliderProcess import TrainSliderProcess
|
||||||
|
from .TrainSliderProcessOld import TrainSliderProcessOld
|
||||||
from .TrainLoRAHack import TrainLoRAHack
|
from .TrainLoRAHack import TrainLoRAHack
|
||||||
from .TrainSDRescaleProcess import TrainSDRescaleProcess
|
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
|
||||||
import torch.nn as nn
|
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):
|
class MeanReduce(nn.Module):
|
||||||
@@ -36,3 +48,147 @@ class Vgg19Critic(nn.Module):
|
|||||||
|
|
||||||
def forward(self, inputs):
|
def forward(self, inputs):
|
||||||
return self.main(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
|
transformers
|
||||||
lycoris_lora
|
lycoris_lora
|
||||||
flatten_json
|
flatten_json
|
||||||
accelerator
|
|
||||||
pyyaml
|
pyyaml
|
||||||
oyaml
|
oyaml
|
||||||
tensorboard
|
tensorboard
|
||||||
@@ -14,4 +13,6 @@ invisible-watermark
|
|||||||
einops
|
einops
|
||||||
accelerate
|
accelerate
|
||||||
toml
|
toml
|
||||||
albumentations
|
albumentations
|
||||||
|
pydantic
|
||||||
|
omegaconf
|
||||||
|
|||||||
12
run.py
12
run.py
@@ -1,5 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
from typing import Union, OrderedDict
|
||||||
|
|
||||||
sys.path.insert(0, os.getcwd())
|
sys.path.insert(0, os.getcwd())
|
||||||
import argparse
|
import argparse
|
||||||
from toolkit.job import get_job
|
from toolkit.job import get_job
|
||||||
@@ -36,6 +38,14 @@ def main():
|
|||||||
action='store_true',
|
action='store_true',
|
||||||
help='Continue running additional jobs even if a job fails'
|
help='Continue running additional jobs even if a job fails'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# flag to continue if failed job
|
||||||
|
parser.add_argument(
|
||||||
|
'-n', '--name',
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||||
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
config_file_list = args.config_file_list
|
config_file_list = args.config_file_list
|
||||||
@@ -49,7 +59,7 @@ def main():
|
|||||||
|
|
||||||
for config_file in config_file_list:
|
for config_file in config_file_list:
|
||||||
try:
|
try:
|
||||||
job = get_job(config_file)
|
job = get_job(config_file, args.name)
|
||||||
job.run()
|
job.run()
|
||||||
job.cleanup()
|
job.cleanup()
|
||||||
jobs_completed += 1
|
jobs_completed += 1
|
||||||
|
|||||||
99
testing/compare_keys.py
Normal file
99
testing/compare_keys.py
Normal file
@@ -0,0 +1,99 @@
|
|||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from diffusers.loaders import LoraLoaderMixin
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
from collections import OrderedDict
|
||||||
|
import json
|
||||||
|
# this was just used to match the vae keys to the diffusers keys
|
||||||
|
# you probably wont need this. Unless they change them.... again... again
|
||||||
|
# on second thought, you probably will
|
||||||
|
|
||||||
|
device = torch.device('cpu')
|
||||||
|
dtype = torch.float32
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
# require at lease one config file
|
||||||
|
parser.add_argument(
|
||||||
|
'file_1',
|
||||||
|
nargs='+',
|
||||||
|
type=str,
|
||||||
|
help='Path to first safe tensor file'
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
'file_2',
|
||||||
|
nargs='+',
|
||||||
|
type=str,
|
||||||
|
help='Path to second safe tensor file'
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
find_matches = False
|
||||||
|
|
||||||
|
state_dict_file_1 = load_file(args.file_1[0])
|
||||||
|
state_dict_1_keys = list(state_dict_file_1.keys())
|
||||||
|
|
||||||
|
state_dict_file_2 = load_file(args.file_2[0])
|
||||||
|
state_dict_2_keys = list(state_dict_file_2.keys())
|
||||||
|
keys_in_both = []
|
||||||
|
|
||||||
|
keys_not_in_state_dict_2 = []
|
||||||
|
for key in state_dict_1_keys:
|
||||||
|
if key not in state_dict_2_keys:
|
||||||
|
keys_not_in_state_dict_2.append(key)
|
||||||
|
|
||||||
|
keys_not_in_state_dict_1 = []
|
||||||
|
for key in state_dict_2_keys:
|
||||||
|
if key not in state_dict_1_keys:
|
||||||
|
keys_not_in_state_dict_1.append(key)
|
||||||
|
|
||||||
|
keys_in_both = []
|
||||||
|
for key in state_dict_1_keys:
|
||||||
|
if key in state_dict_2_keys:
|
||||||
|
keys_in_both.append(key)
|
||||||
|
|
||||||
|
# sort them
|
||||||
|
keys_not_in_state_dict_2.sort()
|
||||||
|
keys_not_in_state_dict_1.sort()
|
||||||
|
keys_in_both.sort()
|
||||||
|
|
||||||
|
|
||||||
|
json_data = {
|
||||||
|
"both": keys_in_both,
|
||||||
|
"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)
|
||||||
|
|
||||||
|
remaining_diffusers_values = OrderedDict()
|
||||||
|
for key in keys_not_in_state_dict_1:
|
||||||
|
remaining_diffusers_values[key] = state_dict_file_2[key]
|
||||||
|
|
||||||
|
# print(remaining_diffusers_values.keys())
|
||||||
|
|
||||||
|
remaining_ldm_values = OrderedDict()
|
||||||
|
for key in keys_not_in_state_dict_2:
|
||||||
|
remaining_ldm_values[key] = state_dict_file_1[key]
|
||||||
|
|
||||||
|
# print(json_data)
|
||||||
|
|
||||||
|
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 os
|
||||||
import json
|
import json
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
import oyaml as yaml
|
import oyaml as yaml
|
||||||
import re
|
import re
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
@@ -15,22 +17,24 @@ def get_cwd_abs_path(path):
|
|||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def preprocess_config(config: OrderedDict):
|
def preprocess_config(config: OrderedDict, name: str = None):
|
||||||
if "job" not in config:
|
if "job" not in config:
|
||||||
raise ValueError("config file must have a job key")
|
raise ValueError("config file must have a job key")
|
||||||
if "config" not in config:
|
if "config" not in config:
|
||||||
raise ValueError("config file must have a config section")
|
raise ValueError("config file must have a config section")
|
||||||
if "name" not in config["config"]:
|
if "name" not in config["config"] and name is None:
|
||||||
raise ValueError("config file must have a config.name key")
|
raise ValueError("config file must have a config.name key")
|
||||||
# we need to replace tags. For now just [name]
|
# we need to replace tags. For now just [name]
|
||||||
name = config["config"]["name"]
|
if name is not None:
|
||||||
|
config["config"]["name"] = name
|
||||||
|
else:
|
||||||
|
name = config["config"]["name"]
|
||||||
config_string = json.dumps(config)
|
config_string = json.dumps(config)
|
||||||
config_string = config_string.replace("[name]", name)
|
config_string = config_string.replace("[name]", name)
|
||||||
config = json.loads(config_string, object_pairs_hook=OrderedDict)
|
config = json.loads(config_string, object_pairs_hook=OrderedDict)
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# Fixes issue where yaml doesnt load exponents correctly
|
# Fixes issue where yaml doesnt load exponents correctly
|
||||||
fixed_loader = yaml.SafeLoader
|
fixed_loader = yaml.SafeLoader
|
||||||
fixed_loader.add_implicit_resolver(
|
fixed_loader.add_implicit_resolver(
|
||||||
@@ -44,7 +48,18 @@ fixed_loader.add_implicit_resolver(
|
|||||||
|\\.(?:nan|NaN|NAN))$''', re.X),
|
|\\.(?:nan|NaN|NAN))$''', re.X),
|
||||||
list(u'-+0123456789.'))
|
list(u'-+0123456789.'))
|
||||||
|
|
||||||
def get_config(config_file_path):
|
|
||||||
|
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
|
# first check if it is in the config folder
|
||||||
config_path = os.path.join(TOOLKIT_ROOT, 'config', config_file_path)
|
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
|
# see if it is in the config folder with any of the possible extensions if it doesnt have one
|
||||||
@@ -67,12 +82,12 @@ def get_config(config_file_path):
|
|||||||
|
|
||||||
# if we found it, check if it is a json or yaml file
|
# 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'):
|
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)
|
config = json.load(f, object_pairs_hook=OrderedDict)
|
||||||
elif real_config_path.endswith('.yaml') or real_config_path.endswith('.yml'):
|
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)
|
config = yaml.load(f, Loader=fixed_loader)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
||||||
|
|
||||||
return preprocess_config(config)
|
return preprocess_config(config, name)
|
||||||
|
|||||||
@@ -1,4 +1,7 @@
|
|||||||
from typing import List
|
import os
|
||||||
|
import time
|
||||||
|
from typing import List, Optional
|
||||||
|
import random
|
||||||
|
|
||||||
|
|
||||||
class SaveConfig:
|
class SaveConfig:
|
||||||
@@ -27,6 +30,7 @@ class SampleConfig:
|
|||||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||||
self.network_multiplier = kwargs.get('network_multiplier', 1)
|
self.network_multiplier = kwargs.get('network_multiplier', 1)
|
||||||
|
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||||
|
|
||||||
|
|
||||||
class NetworkConfig:
|
class NetworkConfig:
|
||||||
@@ -35,13 +39,15 @@ class NetworkConfig:
|
|||||||
rank = kwargs.get('rank', None)
|
rank = kwargs.get('rank', None)
|
||||||
linear = kwargs.get('linear', None)
|
linear = kwargs.get('linear', None)
|
||||||
if rank is not 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
|
self.linear: int = rank
|
||||||
elif linear is not None:
|
elif linear is not None:
|
||||||
self.rank: int = linear
|
self.rank: int = linear
|
||||||
self.linear: int = linear
|
self.linear: int = linear
|
||||||
self.conv: int = kwargs.get('conv', None)
|
self.conv: int = kwargs.get('conv', None)
|
||||||
self.alpha: float = kwargs.get('alpha', 1.0)
|
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:
|
class TrainConfig:
|
||||||
@@ -60,7 +66,7 @@ class TrainConfig:
|
|||||||
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
||||||
self.optimizer_params = kwargs.get('optimizer_params', {})
|
self.optimizer_params = kwargs.get('optimizer_params', {})
|
||||||
self.skip_first_sample = kwargs.get('skip_first_sample', False)
|
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:
|
class ModelConfig:
|
||||||
@@ -69,6 +75,8 @@ class ModelConfig:
|
|||||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||||
self.is_v_pred: bool = kwargs.get('is_v_pred', 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:
|
if self.name_or_path is None:
|
||||||
raise ValueError('name_or_path must be specified')
|
raise ValueError('name_or_path must be specified')
|
||||||
@@ -99,3 +107,200 @@ class SliderConfig:
|
|||||||
anchors = [SliderConfigAnchors(**anchor) for anchor in anchors]
|
anchors = [SliderConfigAnchors(**anchor) for anchor in anchors]
|
||||||
self.anchors: List[SliderConfigAnchors] = anchors
|
self.anchors: List[SliderConfigAnchors] = anchors
|
||||||
self.resolutions: List[List[int]] = kwargs.get('resolutions', [[512, 512]])
|
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 os
|
||||||
import random
|
import random
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from PIL.ImageOps import exif_transpose
|
from PIL.ImageOps import exif_transpose
|
||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
import albumentations as A
|
||||||
|
|
||||||
|
|
||||||
class ImageDataset(Dataset):
|
class ImageDataset(Dataset):
|
||||||
@@ -38,7 +42,7 @@ class ImageDataset(Dataset):
|
|||||||
|
|
||||||
self.transform = transforms.Compose([
|
self.transform = transforms.Compose([
|
||||||
transforms.ToTensor(),
|
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):
|
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 self.random_scale and min_img_size > self.resolution:
|
||||||
if min_img_size < self.resolution:
|
if min_img_size < self.resolution:
|
||||||
print(
|
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
|
scale_size = self.resolution
|
||||||
else:
|
else:
|
||||||
scale_size = random.randint(self.resolution, int(min_img_size))
|
scale_size = random.randint(self.resolution, int(min_img_size))
|
||||||
@@ -78,3 +82,124 @@ class ImageDataset(Dataset):
|
|||||||
img = self.transform(img)
|
img = self.transform(img)
|
||||||
|
|
||||||
return 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,8 +1,13 @@
|
|||||||
|
from typing import Union, OrderedDict
|
||||||
|
|
||||||
from toolkit.config import get_config
|
from toolkit.config import get_config
|
||||||
|
|
||||||
|
|
||||||
def get_job(config_path):
|
def get_job(
|
||||||
config = get_config(config_path)
|
config_path: Union[str, dict, OrderedDict],
|
||||||
|
name=None
|
||||||
|
):
|
||||||
|
config = get_config(config_path, name)
|
||||||
if not config['job']:
|
if not config['job']:
|
||||||
raise ValueError('config file is invalid. Missing "job" key')
|
raise ValueError('config file is invalid. Missing "job" key')
|
||||||
|
|
||||||
@@ -13,9 +18,27 @@ def get_job(config_path):
|
|||||||
if job == 'train':
|
if job == 'train':
|
||||||
from jobs import TrainJob
|
from jobs import TrainJob
|
||||||
return TrainJob(config)
|
return TrainJob(config)
|
||||||
|
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':
|
# elif job == 'train':
|
||||||
# from jobs import TrainJob
|
# from jobs import TrainJob
|
||||||
# return TrainJob(config)
|
# return TrainJob(config)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f'Unknown job type {job}')
|
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:
|
for key in keys:
|
||||||
if key.startswith("cond_stage_model.transformer"):
|
if key.startswith("cond_stage_model.transformer"):
|
||||||
text_model_dict[key[len("cond_stage_model.transformer."):]] = checkpoint[key]
|
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
|
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)
|
text_model = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14").to(device)
|
||||||
logging.set_verbosity_warning()
|
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)
|
info = text_model.load_state_dict(converted_text_encoder_checkpoint)
|
||||||
print("loading text encoder:", info)
|
print("loading text encoder:", info)
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
from torch.utils.checkpoint import checkpoint
|
||||||
|
|
||||||
|
|
||||||
class ReductionKernel(nn.Module):
|
class ReductionKernel(nn.Module):
|
||||||
@@ -29,3 +30,15 @@ class ReductionKernel(nn.Module):
|
|||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
return nn.functional.conv2d(x, self.kernel, stride=self.kernel_size, padding=0, groups=1)
|
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)
|
||||||
|
|||||||
@@ -6,12 +6,14 @@
|
|||||||
import os
|
import os
|
||||||
import math
|
import math
|
||||||
from typing import Optional, List, Type, Set, Literal
|
from typing import Optional, List, Type, Set, Literal
|
||||||
|
from collections import OrderedDict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from diffusers import UNet2DConditionModel
|
from diffusers import UNet2DConditionModel
|
||||||
from safetensors.torch import save_file
|
from safetensors.torch import save_file
|
||||||
|
|
||||||
|
from toolkit.metadata import add_model_hash_to_meta
|
||||||
|
|
||||||
UNET_TARGET_REPLACE_MODULE_TRANSFORMER = [
|
UNET_TARGET_REPLACE_MODULE_TRANSFORMER = [
|
||||||
"Transformer2DModel", # どうやらこっちの方らしい? # attn1, 2
|
"Transformer2DModel", # どうやらこっちの方らしい? # attn1, 2
|
||||||
@@ -31,7 +33,7 @@ TRAINING_METHODS = Literal[
|
|||||||
"innoxattn", # train all layers except self attention layers
|
"innoxattn", # train all layers except self attention layers
|
||||||
"selfattn", # ESD-u, train only self attention layers
|
"selfattn", # ESD-u, train only self attention layers
|
||||||
"xattn", # ESD-x, train only x attention layers
|
"xattn", # ESD-x, train only x attention layers
|
||||||
"full", # train all layers
|
"full", # train all layers
|
||||||
# "notime",
|
# "notime",
|
||||||
# "xlayer",
|
# "xlayer",
|
||||||
# "outxattn",
|
# "outxattn",
|
||||||
@@ -48,12 +50,12 @@ class LoRAModule(nn.Module):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
lora_name,
|
lora_name,
|
||||||
org_module: nn.Module,
|
org_module: nn.Module,
|
||||||
multiplier=1.0,
|
multiplier=1.0,
|
||||||
lora_dim=4,
|
lora_dim=4,
|
||||||
alpha=1,
|
alpha=1,
|
||||||
):
|
):
|
||||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -102,19 +104,19 @@ class LoRAModule(nn.Module):
|
|||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
return (
|
return (
|
||||||
self.org_forward(x)
|
self.org_forward(x)
|
||||||
+ self.lora_up(self.lora_down(x)) * self.multiplier * self.scale
|
+ self.lora_up(self.lora_down(x)) * self.multiplier * self.scale
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class LoRANetwork(nn.Module):
|
class LoRANetwork(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
unet: UNet2DConditionModel,
|
unet: UNet2DConditionModel,
|
||||||
rank: int = 4,
|
rank: int = 4,
|
||||||
multiplier: float = 1.0,
|
multiplier: float = 1.0,
|
||||||
alpha: float = 1.0,
|
alpha: float = 1.0,
|
||||||
train_method: TRAINING_METHODS = "full",
|
train_method: TRAINING_METHODS = "full",
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -140,7 +142,7 @@ class LoRANetwork(nn.Module):
|
|||||||
lora_names = set()
|
lora_names = set()
|
||||||
for lora in self.unet_loras:
|
for lora in self.unet_loras:
|
||||||
assert (
|
assert (
|
||||||
lora.lora_name not in lora_names
|
lora.lora_name not in lora_names
|
||||||
), f"duplicated lora name: {lora.lora_name}. {lora_names}"
|
), f"duplicated lora name: {lora.lora_name}. {lora_names}"
|
||||||
lora_names.add(lora.lora_name)
|
lora_names.add(lora.lora_name)
|
||||||
|
|
||||||
@@ -157,13 +159,13 @@ class LoRANetwork(nn.Module):
|
|||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
def create_modules(
|
def create_modules(
|
||||||
self,
|
self,
|
||||||
prefix: str,
|
prefix: str,
|
||||||
root_module: nn.Module,
|
root_module: nn.Module,
|
||||||
target_replace_modules: List[str],
|
target_replace_modules: List[str],
|
||||||
rank: int,
|
rank: int,
|
||||||
multiplier: float,
|
multiplier: float,
|
||||||
train_method: TRAINING_METHODS,
|
train_method: TRAINING_METHODS,
|
||||||
) -> list:
|
) -> list:
|
||||||
loras = []
|
loras = []
|
||||||
|
|
||||||
@@ -212,6 +214,8 @@ class LoRANetwork(nn.Module):
|
|||||||
|
|
||||||
def save_weights(self, file, dtype=None, metadata: Optional[dict] = None):
|
def save_weights(self, file, dtype=None, metadata: Optional[dict] = None):
|
||||||
state_dict = self.state_dict()
|
state_dict = self.state_dict()
|
||||||
|
if metadata is None:
|
||||||
|
metadata = OrderedDict()
|
||||||
|
|
||||||
if dtype is not None:
|
if dtype is not None:
|
||||||
for key in list(state_dict.keys()):
|
for key in list(state_dict.keys()):
|
||||||
@@ -221,9 +225,10 @@ class LoRANetwork(nn.Module):
|
|||||||
|
|
||||||
for key in list(state_dict.keys()):
|
for key in list(state_dict.keys()):
|
||||||
if not key.startswith("lora"):
|
if not key.startswith("lora"):
|
||||||
# lora以外除外
|
# remove any not lora
|
||||||
del state_dict[key]
|
del state_dict[key]
|
||||||
|
|
||||||
|
metadata = add_model_hash_to_meta(state_dict, metadata)
|
||||||
if os.path.splitext(file)[1] == ".safetensors":
|
if os.path.splitext(file)[1] == ".safetensors":
|
||||||
save_file(state_dict, file, metadata)
|
save_file(state_dict, file, metadata)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
|
import math
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
import sys
|
import sys
|
||||||
from typing import List, Optional, Dict, Type, Union
|
from typing import List, Optional, Dict, Type, Union
|
||||||
|
|
||||||
@@ -9,7 +11,174 @@ from .paths import SD_SCRIPTS_ROOT
|
|||||||
|
|
||||||
sys.path.append(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):
|
class LoRASpecialNetwork(LoRANetwork):
|
||||||
@@ -70,6 +239,7 @@ class LoRASpecialNetwork(LoRANetwork):
|
|||||||
self.dropout = dropout
|
self.dropout = dropout
|
||||||
self.rank_dropout = rank_dropout
|
self.rank_dropout = rank_dropout
|
||||||
self.module_dropout = module_dropout
|
self.module_dropout = module_dropout
|
||||||
|
self.is_checkpointing = False
|
||||||
|
|
||||||
if modules_dim is not None:
|
if modules_dim is not None:
|
||||||
print(f"create LoRA network from weights")
|
print(f"create LoRA network from weights")
|
||||||
@@ -236,11 +406,11 @@ class LoRASpecialNetwork(LoRANetwork):
|
|||||||
torch.save(state_dict, file)
|
torch.save(state_dict, file)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def multiplier(self):
|
def multiplier(self) -> Union[float, List[float]]:
|
||||||
return self._multiplier
|
return self._multiplier
|
||||||
|
|
||||||
@multiplier.setter
|
@multiplier.setter
|
||||||
def multiplier(self, value):
|
def multiplier(self, value: Union[float, List[float]]):
|
||||||
self._multiplier = value
|
self._multiplier = value
|
||||||
self._update_lora_multiplier()
|
self._update_lora_multiplier()
|
||||||
|
|
||||||
@@ -261,6 +431,8 @@ class LoRASpecialNetwork(LoRANetwork):
|
|||||||
for lora in self.text_encoder_loras:
|
for lora in self.text_encoder_loras:
|
||||||
lora.multiplier = 0
|
lora.multiplier = 0
|
||||||
|
|
||||||
|
# called when the context manager is entered
|
||||||
|
# ie: with network:
|
||||||
def __enter__(self):
|
def __enter__(self):
|
||||||
self.is_active = True
|
self.is_active = True
|
||||||
self._update_lora_multiplier()
|
self._update_lora_multiplier()
|
||||||
@@ -278,3 +450,29 @@ class LoRASpecialNetwork(LoRANetwork):
|
|||||||
loras += self.text_encoder_loras
|
loras += self.text_encoder_loras
|
||||||
for lora in loras:
|
for lora in loras:
|
||||||
lora.to(device, dtype)
|
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()
|
||||||
|
|||||||
File diff suppressed because one or more lines are too long
@@ -1,18 +1,23 @@
|
|||||||
import json
|
import json
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
|
import safetensors
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
|
|
||||||
from info import software_meta
|
from info import software_meta
|
||||||
|
from toolkit.train_tools import addnet_hash_legacy
|
||||||
|
from toolkit.train_tools import addnet_hash_safetensors
|
||||||
|
|
||||||
|
|
||||||
def get_meta_for_safetensors(meta: OrderedDict, name=None) -> OrderedDict:
|
def get_meta_for_safetensors(meta: OrderedDict, name=None, add_software_info=True) -> OrderedDict:
|
||||||
# stringify the meta and reparse OrderedDict to replace [name] with name
|
# stringify the meta and reparse OrderedDict to replace [name] with name
|
||||||
meta_string = json.dumps(meta)
|
meta_string = json.dumps(meta)
|
||||||
if name is not None:
|
if name is not None:
|
||||||
meta_string = meta_string.replace("[name]", name)
|
meta_string = meta_string.replace("[name]", name)
|
||||||
save_meta = json.loads(meta_string, object_pairs_hook=OrderedDict)
|
save_meta = json.loads(meta_string, object_pairs_hook=OrderedDict)
|
||||||
save_meta["software"] = software_meta
|
if add_software_info:
|
||||||
|
save_meta["software"] = software_meta
|
||||||
# safetensors can only be one level deep
|
# safetensors can only be one level deep
|
||||||
for key, value in save_meta.items():
|
for key, value in save_meta.items():
|
||||||
# if not float, int, bool, or str, convert to json string
|
# if not float, int, bool, or str, convert to json string
|
||||||
@@ -21,6 +26,46 @@ def get_meta_for_safetensors(meta: OrderedDict, name=None) -> OrderedDict:
|
|||||||
return save_meta
|
return save_meta
|
||||||
|
|
||||||
|
|
||||||
|
def add_model_hash_to_meta(state_dict, meta: OrderedDict) -> OrderedDict:
|
||||||
|
"""Precalculate the model hashes needed by sd-webui-additional-networks to
|
||||||
|
save time on indexing the model later."""
|
||||||
|
|
||||||
|
# Because writing user metadata to the file can change the result of
|
||||||
|
# sd_models.model_hash(), only retain the training metadata for purposes of
|
||||||
|
# calculating the hash, as they are meant to be immutable
|
||||||
|
metadata = {k: v for k, v in meta.items() if k.startswith("ss_")}
|
||||||
|
|
||||||
|
bytes = safetensors.torch.save(state_dict, metadata)
|
||||||
|
b = BytesIO(bytes)
|
||||||
|
|
||||||
|
model_hash = addnet_hash_safetensors(b)
|
||||||
|
legacy_hash = addnet_hash_legacy(b)
|
||||||
|
meta["sshs_model_hash"] = model_hash
|
||||||
|
meta["sshs_legacy_hash"] = legacy_hash
|
||||||
|
return meta
|
||||||
|
|
||||||
|
|
||||||
|
def add_base_model_info_to_meta(
|
||||||
|
meta: OrderedDict,
|
||||||
|
base_model: str = None,
|
||||||
|
is_v1: bool = False,
|
||||||
|
is_v2: bool = False,
|
||||||
|
is_xl: bool = False,
|
||||||
|
) -> OrderedDict:
|
||||||
|
if base_model is not None:
|
||||||
|
meta['ss_base_model'] = base_model
|
||||||
|
elif is_v2:
|
||||||
|
meta['ss_v2'] = True
|
||||||
|
meta['ss_base_model_version'] = 'sd_2.1'
|
||||||
|
|
||||||
|
elif is_xl:
|
||||||
|
meta['ss_base_model_version'] = 'sdxl_1.0'
|
||||||
|
else:
|
||||||
|
# default to v1.5
|
||||||
|
meta['ss_base_model_version'] = 'sd_1.5'
|
||||||
|
return meta
|
||||||
|
|
||||||
|
|
||||||
def parse_metadata_from_safetensors(meta: OrderedDict) -> OrderedDict:
|
def parse_metadata_from_safetensors(meta: OrderedDict) -> OrderedDict:
|
||||||
parsed_meta = OrderedDict()
|
parsed_meta = OrderedDict()
|
||||||
for key, value in meta.items():
|
for key, value in meta.items():
|
||||||
|
|||||||
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)
|
||||||
@@ -27,6 +27,17 @@ def get_optimizer(
|
|||||||
optimizer = dadaptation.DAdaptAdam(params, lr=use_lr, **optimizer_params)
|
optimizer = dadaptation.DAdaptAdam(params, lr=use_lr, **optimizer_params)
|
||||||
# warn user that dadaptation is deprecated
|
# warn user that dadaptation is deprecated
|
||||||
print("WARNING: Dadaptation optimizer type has been changed to DadaptationAdam. Please update your config.")
|
print("WARNING: Dadaptation optimizer type has been changed to DadaptationAdam. Please update your config.")
|
||||||
|
elif lower_type.startswith("prodigy"):
|
||||||
|
from prodigyopt import Prodigy
|
||||||
|
|
||||||
|
print("Using Prodigy optimizer")
|
||||||
|
use_lr = learning_rate
|
||||||
|
if use_lr < 0.1:
|
||||||
|
# dadaptation uses different lr that is values of 0.1 to 1.0. default to 1.0
|
||||||
|
use_lr = 1.0
|
||||||
|
# let net be the neural network you want to train
|
||||||
|
# you can choose weight decay value based on your problem, 0 by default
|
||||||
|
optimizer = Prodigy(params, lr=use_lr, **optimizer_params)
|
||||||
elif lower_type.endswith("8bit"):
|
elif lower_type.endswith("8bit"):
|
||||||
import bitsandbytes
|
import bitsandbytes
|
||||||
|
|
||||||
@@ -43,6 +54,8 @@ def get_optimizer(
|
|||||||
elif lower_type == 'lion':
|
elif lower_type == 'lion':
|
||||||
from lion_pytorch import Lion
|
from lion_pytorch import Lion
|
||||||
return Lion(params, lr=learning_rate, **optimizer_params)
|
return Lion(params, lr=learning_rate, **optimizer_params)
|
||||||
|
elif lower_type == 'adagrad':
|
||||||
|
optimizer = torch.optim.Adagrad(params, lr=float(learning_rate), **optimizer_params)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
||||||
return optimizer
|
return optimizer
|
||||||
|
|||||||
@@ -5,6 +5,12 @@ CONFIG_ROOT = os.path.join(TOOLKIT_ROOT, 'config')
|
|||||||
SD_SCRIPTS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories", "sd-scripts")
|
SD_SCRIPTS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories", "sd-scripts")
|
||||||
REPOS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories")
|
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):
|
def get_path(path):
|
||||||
# we allow absolute paths, but if it is not absolute, we assume it is relative to the toolkit root
|
# we allow absolute paths, but if it is not absolute, we assume it is relative to the toolkit root
|
||||||
|
|||||||
601
toolkit/pipelines.py
Normal file
601
toolkit/pipelines.py
Normal file
@@ -0,0 +1,601 @@
|
|||||||
|
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):
|
||||||
|
# super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
def predict_noise(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]] = None,
|
||||||
|
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_inference_steps: int = 50,
|
||||||
|
guidance_scale: float = 5.0,
|
||||||
|
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||||
|
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_images_per_prompt: Optional[int] = 1,
|
||||||
|
eta: float = 0.0,
|
||||||
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
|
latents: Optional[torch.FloatTensor] = None,
|
||||||
|
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
guidance_rescale: float = 0.0,
|
||||||
|
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||||
|
timestep: Optional[int] = None,
|
||||||
|
):
|
||||||
|
r"""
|
||||||
|
Function invoked when calling the pipeline for generation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||||
|
instead.
|
||||||
|
prompt_2 (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||||
|
used in both text-encoders
|
||||||
|
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||||
|
The height in pixels of the generated image.
|
||||||
|
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||||
|
The width in pixels of the generated image.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||||
|
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||||
|
expense of slower inference.
|
||||||
|
denoising_end (`float`, *optional*):
|
||||||
|
When specified, determines the fraction (between 0.0 and 1.0) of the total denoising process to be
|
||||||
|
completed before it is intentionally prematurely terminated. As a result, the returned sample will
|
||||||
|
still retain a substantial amount of noise as determined by the discrete timesteps selected by the
|
||||||
|
scheduler. The denoising_end parameter should ideally be utilized when this pipeline forms a part of a
|
||||||
|
"Mixture of Denoisers" multi-pipeline setup, as elaborated in [**Refining the Image
|
||||||
|
Output**](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl#refining-the-image-output)
|
||||||
|
guidance_scale (`float`, *optional*, defaults to 7.5):
|
||||||
|
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||||
|
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||||
|
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||||
|
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||||
|
usually at the expense of lower image quality.
|
||||||
|
negative_prompt (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||||
|
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||||
|
less than `1`).
|
||||||
|
negative_prompt_2 (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
|
||||||
|
`text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders
|
||||||
|
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||||
|
The number of images to generate per prompt.
|
||||||
|
eta (`float`, *optional*, defaults to 0.0):
|
||||||
|
Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to
|
||||||
|
[`schedulers.DDIMScheduler`], will be ignored for others.
|
||||||
|
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||||
|
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||||
|
to make generation deterministic.
|
||||||
|
latents (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||||
|
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||||
|
tensor will ge generated by sampling using the supplied random `generator`.
|
||||||
|
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||||
|
provided, text embeddings will be generated from `prompt` input argument.
|
||||||
|
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||||
|
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||||
|
argument.
|
||||||
|
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||||
|
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||||
|
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||||
|
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
|
||||||
|
input argument.
|
||||||
|
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||||
|
The output format of the generate image. Choose between
|
||||||
|
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||||
|
return_dict (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead
|
||||||
|
of a plain tuple.
|
||||||
|
callback (`Callable`, *optional*):
|
||||||
|
A function that will be called every `callback_steps` steps during inference. The function will be
|
||||||
|
called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||||
|
callback_steps (`int`, *optional*, defaults to 1):
|
||||||
|
The frequency at which the `callback` function will be called. If not specified, the callback will be
|
||||||
|
called at every step.
|
||||||
|
cross_attention_kwargs (`dict`, *optional*):
|
||||||
|
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||||
|
`self.processor` in
|
||||||
|
[diffusers.cross_attention](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/cross_attention.py).
|
||||||
|
guidance_rescale (`float`, *optional*, defaults to 0.7):
|
||||||
|
Guidance rescale factor proposed by [Common Diffusion Noise Schedules and Sample Steps are
|
||||||
|
Flawed](https://arxiv.org/pdf/2305.08891.pdf) `guidance_scale` is defined as `φ` in equation 16. of
|
||||||
|
[Common Diffusion Noise Schedules and Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf).
|
||||||
|
Guidance rescale factor should fix overexposure when using zero terminal SNR.
|
||||||
|
original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||||
|
If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled.
|
||||||
|
`original_size` defaults to `(width, height)` if not specified. Part of SDXL's micro-conditioning as
|
||||||
|
explained in section 2.2 of
|
||||||
|
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
|
||||||
|
`crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position
|
||||||
|
`crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting
|
||||||
|
`crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of
|
||||||
|
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||||
|
For most cases, `target_size` should be set to the desired height and width of the generated image. If
|
||||||
|
not specified it will default to `(width, height)`. Part of SDXL's micro-conditioning as explained in
|
||||||
|
section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] or `tuple`:
|
||||||
|
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] if `return_dict` is True, otherwise a
|
||||||
|
`tuple`. When returning a tuple, the first element is a list with the generated images.
|
||||||
|
"""
|
||||||
|
# if not predict_noise:
|
||||||
|
# # call parent
|
||||||
|
# return super().__call__(
|
||||||
|
# prompt=prompt,
|
||||||
|
# prompt_2=prompt_2,
|
||||||
|
# height=height,
|
||||||
|
# width=width,
|
||||||
|
# num_inference_steps=num_inference_steps,
|
||||||
|
# denoising_end=denoising_end,
|
||||||
|
# guidance_scale=guidance_scale,
|
||||||
|
# negative_prompt=negative_prompt,
|
||||||
|
# negative_prompt_2=negative_prompt_2,
|
||||||
|
# num_images_per_prompt=num_images_per_prompt,
|
||||||
|
# eta=eta,
|
||||||
|
# generator=generator,
|
||||||
|
# latents=latents,
|
||||||
|
# prompt_embeds=prompt_embeds,
|
||||||
|
# negative_prompt_embeds=negative_prompt_embeds,
|
||||||
|
# pooled_prompt_embeds=pooled_prompt_embeds,
|
||||||
|
# negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||||
|
# output_type=output_type,
|
||||||
|
# return_dict=return_dict,
|
||||||
|
# callback=callback,
|
||||||
|
# callback_steps=callback_steps,
|
||||||
|
# cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
# guidance_rescale=guidance_rescale,
|
||||||
|
# original_size=original_size,
|
||||||
|
# crops_coords_top_left=crops_coords_top_left,
|
||||||
|
# target_size=target_size,
|
||||||
|
# )
|
||||||
|
|
||||||
|
# 0. Default height and width to unet
|
||||||
|
height = self.default_sample_size * self.vae_scale_factor
|
||||||
|
width = self.default_sample_size * self.vae_scale_factor
|
||||||
|
|
||||||
|
original_size = (height, width)
|
||||||
|
target_size = (height, width)
|
||||||
|
|
||||||
|
# 2. Define call parameters
|
||||||
|
if prompt is not None and isinstance(prompt, str):
|
||||||
|
batch_size = 1
|
||||||
|
elif prompt is not None and isinstance(prompt, list):
|
||||||
|
batch_size = len(prompt)
|
||||||
|
else:
|
||||||
|
batch_size = prompt_embeds.shape[0]
|
||||||
|
|
||||||
|
device = self._execution_device
|
||||||
|
|
||||||
|
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||||
|
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||||
|
# corresponds to doing no classifier free guidance.
|
||||||
|
do_classifier_free_guidance = guidance_scale > 1.0
|
||||||
|
|
||||||
|
# 3. Encode input prompt
|
||||||
|
text_encoder_lora_scale = (
|
||||||
|
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
|
||||||
|
)
|
||||||
|
(
|
||||||
|
prompt_embeds,
|
||||||
|
negative_prompt_embeds,
|
||||||
|
pooled_prompt_embeds,
|
||||||
|
negative_pooled_prompt_embeds,
|
||||||
|
) = self.encode_prompt(
|
||||||
|
prompt=prompt,
|
||||||
|
prompt_2=prompt_2,
|
||||||
|
device=device,
|
||||||
|
num_images_per_prompt=num_images_per_prompt,
|
||||||
|
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||||
|
negative_prompt=negative_prompt,
|
||||||
|
negative_prompt_2=negative_prompt_2,
|
||||||
|
prompt_embeds=prompt_embeds,
|
||||||
|
negative_prompt_embeds=negative_prompt_embeds,
|
||||||
|
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||||
|
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||||
|
lora_scale=text_encoder_lora_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Prepare timesteps
|
||||||
|
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||||
|
|
||||||
|
# 5. Prepare latent variables
|
||||||
|
num_channels_latents = self.unet.config.in_channels
|
||||||
|
latents = self.prepare_latents(
|
||||||
|
batch_size * num_images_per_prompt,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
prompt_embeds.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 7. Prepare added time ids & embeddings
|
||||||
|
add_text_embeds = pooled_prompt_embeds
|
||||||
|
add_time_ids = self._get_add_time_ids(
|
||||||
|
original_size, crops_coords_top_left, target_size, dtype=prompt_embeds.dtype
|
||||||
|
).to(device) # TODO DOES NOT CAST ORIGINALLY
|
||||||
|
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||||
|
add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0)
|
||||||
|
add_time_ids = torch.cat([add_time_ids, add_time_ids], dim=0)
|
||||||
|
|
||||||
|
prompt_embeds = prompt_embeds.to(device)
|
||||||
|
add_text_embeds = add_text_embeds.to(device)
|
||||||
|
add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1)
|
||||||
|
|
||||||
|
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||||
|
|
||||||
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, timestep)
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids}
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input,
|
||||||
|
timestep,
|
||||||
|
encoder_hidden_states=prompt_embeds,
|
||||||
|
cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
added_cond_kwargs=added_cond_kwargs,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||||
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||||
|
|
||||||
|
if do_classifier_free_guidance and 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)
|
||||||
|
|
||||||
|
return noise_pred
|
||||||
|
|
||||||
|
def enable_model_cpu_offload(self, gpu_id=0):
|
||||||
|
print('Called cpu offload', gpu_id)
|
||||||
|
# fuck off
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class CustomStableDiffusionPipeline(StableDiffusionPipeline):
|
||||||
|
|
||||||
|
# replace the call so it matches SDXL call so we can use the same code and also stop early
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]] = None,
|
||||||
|
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
height: Optional[int] = None,
|
||||||
|
width: Optional[int] = None,
|
||||||
|
num_inference_steps: int = 50,
|
||||||
|
denoising_end: Optional[float] = None,
|
||||||
|
guidance_scale: float = 5.0,
|
||||||
|
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||||
|
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_images_per_prompt: Optional[int] = 1,
|
||||||
|
eta: float = 0.0,
|
||||||
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
|
latents: Optional[torch.FloatTensor] = None,
|
||||||
|
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
output_type: Optional[str] = "pil",
|
||||||
|
return_dict: bool = True,
|
||||||
|
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||||
|
callback_steps: int = 1,
|
||||||
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
guidance_rescale: float = 0.0,
|
||||||
|
original_size: Optional[Tuple[int, int]] = None,
|
||||||
|
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||||
|
target_size: Optional[Tuple[int, int]] = None,
|
||||||
|
):
|
||||||
|
# 0. Default height and width to unet
|
||||||
|
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
|
||||||
|
# 1. Check inputs. Raise error if not correct
|
||||||
|
self.check_inputs(
|
||||||
|
prompt, height, width, callback_steps, negative_prompt, prompt_embeds, negative_prompt_embeds
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Define call parameters
|
||||||
|
if prompt is not None and isinstance(prompt, str):
|
||||||
|
batch_size = 1
|
||||||
|
elif prompt is not None and isinstance(prompt, list):
|
||||||
|
batch_size = len(prompt)
|
||||||
|
else:
|
||||||
|
batch_size = prompt_embeds.shape[0]
|
||||||
|
|
||||||
|
device = self._execution_device
|
||||||
|
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||||
|
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||||
|
# corresponds to doing no classifier free guidance.
|
||||||
|
do_classifier_free_guidance = guidance_scale > 1.0
|
||||||
|
|
||||||
|
# 3. Encode input prompt
|
||||||
|
text_encoder_lora_scale = (
|
||||||
|
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
|
||||||
|
)
|
||||||
|
prompt_embeds = self._encode_prompt(
|
||||||
|
prompt,
|
||||||
|
device,
|
||||||
|
num_images_per_prompt,
|
||||||
|
do_classifier_free_guidance,
|
||||||
|
negative_prompt,
|
||||||
|
prompt_embeds=prompt_embeds,
|
||||||
|
negative_prompt_embeds=negative_prompt_embeds,
|
||||||
|
lora_scale=text_encoder_lora_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Prepare timesteps
|
||||||
|
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||||
|
timesteps = self.scheduler.timesteps
|
||||||
|
|
||||||
|
# 5. Prepare latent variables
|
||||||
|
num_channels_latents = self.unet.config.in_channels
|
||||||
|
latents = self.prepare_latents(
|
||||||
|
batch_size * num_images_per_prompt,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
prompt_embeds.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||||
|
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||||
|
|
||||||
|
# 7. Denoising loop
|
||||||
|
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||||
|
|
||||||
|
# 7.1 Apply denoising_end
|
||||||
|
if denoising_end is not None and type(denoising_end) == float and denoising_end > 0 and denoising_end < 1:
|
||||||
|
discrete_timestep_cutoff = int(
|
||||||
|
round(
|
||||||
|
self.scheduler.config.num_train_timesteps
|
||||||
|
- (denoising_end * self.scheduler.config.num_train_timesteps)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
num_inference_steps = len(list(filter(lambda ts: ts >= discrete_timestep_cutoff, timesteps)))
|
||||||
|
timesteps = timesteps[:num_inference_steps]
|
||||||
|
|
||||||
|
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||||
|
for i, t in enumerate(timesteps):
|
||||||
|
# expand the latents if we are doing classifier free guidance
|
||||||
|
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||||
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input,
|
||||||
|
t,
|
||||||
|
encoder_hidden_states=prompt_embeds,
|
||||||
|
cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||||
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||||
|
|
||||||
|
if do_classifier_free_guidance and 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)
|
||||||
|
|
||||||
|
# compute the previous noisy sample x_t -> x_t-1
|
||||||
|
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||||
|
|
||||||
|
# call the callback, if provided
|
||||||
|
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||||
|
progress_bar.update()
|
||||||
|
if callback is not None and i % callback_steps == 0:
|
||||||
|
callback(i, t, latents)
|
||||||
|
|
||||||
|
if not output_type == "latent":
|
||||||
|
image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0]
|
||||||
|
image, has_nsfw_concept = self.run_safety_checker(image, device, prompt_embeds.dtype)
|
||||||
|
else:
|
||||||
|
image = latents
|
||||||
|
has_nsfw_concept = None
|
||||||
|
|
||||||
|
if has_nsfw_concept is None:
|
||||||
|
do_denormalize = [True] * image.shape[0]
|
||||||
|
else:
|
||||||
|
do_denormalize = [not has_nsfw for has_nsfw in has_nsfw_concept]
|
||||||
|
|
||||||
|
image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize)
|
||||||
|
|
||||||
|
# Offload last model to CPU
|
||||||
|
if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None:
|
||||||
|
self.final_offload_hook.offload()
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (image, has_nsfw_concept)
|
||||||
|
|
||||||
|
return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)
|
||||||
|
|
||||||
|
# some of the inputs are to keep it compatible with sdx
|
||||||
|
def predict_noise(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]] = None,
|
||||||
|
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_inference_steps: int = 50,
|
||||||
|
guidance_scale: float = 5.0,
|
||||||
|
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||||
|
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_images_per_prompt: Optional[int] = 1,
|
||||||
|
eta: float = 0.0,
|
||||||
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
|
latents: Optional[torch.FloatTensor] = None,
|
||||||
|
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
guidance_rescale: float = 0.0,
|
||||||
|
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||||
|
timestep: Optional[int] = None,
|
||||||
|
):
|
||||||
|
|
||||||
|
# 0. Default height and width to unet
|
||||||
|
height = self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
width = self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
|
||||||
|
# 2. Define call parameters
|
||||||
|
if prompt is not None and isinstance(prompt, str):
|
||||||
|
batch_size = 1
|
||||||
|
elif prompt is not None and isinstance(prompt, list):
|
||||||
|
batch_size = len(prompt)
|
||||||
|
else:
|
||||||
|
batch_size = prompt_embeds.shape[0]
|
||||||
|
|
||||||
|
device = self._execution_device
|
||||||
|
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||||
|
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||||
|
# corresponds to doing no classifier free guidance.
|
||||||
|
do_classifier_free_guidance = guidance_scale > 1.0
|
||||||
|
|
||||||
|
# 3. Encode input prompt
|
||||||
|
text_encoder_lora_scale = (
|
||||||
|
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
|
||||||
|
)
|
||||||
|
prompt_embeds = self._encode_prompt(
|
||||||
|
prompt,
|
||||||
|
device,
|
||||||
|
num_images_per_prompt,
|
||||||
|
do_classifier_free_guidance,
|
||||||
|
negative_prompt,
|
||||||
|
prompt_embeds=prompt_embeds,
|
||||||
|
negative_prompt_embeds=negative_prompt_embeds,
|
||||||
|
lora_scale=text_encoder_lora_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Prepare timesteps
|
||||||
|
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||||
|
|
||||||
|
# 5. Prepare latent variables
|
||||||
|
num_channels_latents = self.unet.config.in_channels
|
||||||
|
latents = self.prepare_latents(
|
||||||
|
batch_size * num_images_per_prompt,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
prompt_embeds.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
|
||||||
|
# expand the latents if we are doing classifier free guidance
|
||||||
|
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||||
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, timestep)
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input,
|
||||||
|
timestep,
|
||||||
|
encoder_hidden_states=prompt_embeds,
|
||||||
|
cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||||
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||||
|
|
||||||
|
if do_classifier_free_guidance and 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)
|
||||||
|
|
||||||
|
return noise_pred
|
||||||
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,25 +1,69 @@
|
|||||||
from typing import Union, OrderedDict
|
import gc
|
||||||
|
import typing
|
||||||
|
from typing import Union, OrderedDict, List, Tuple
|
||||||
import sys
|
import sys
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||||
from safetensors.torch import save_file
|
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 library.sdxl_train_util import _load_target_model as load_sdxl_target_model
|
||||||
from toolkit.train_tools import get_torch_dtype
|
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(REPOS_ROOT)
|
||||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||||
|
sys.path.append(os.path.join(REPOS_ROOT, 'sd-scripts'))
|
||||||
from leco import train_util
|
from leco import train_util
|
||||||
import torch
|
import torch
|
||||||
from library import model_util
|
from library.model_util import convert_unet_state_dict_to_sd, convert_text_encoder_state_dict_to_sd_v2
|
||||||
from library.sdxl_model_util import convert_text_encoder_2_state_dict_to_sdxl
|
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:
|
class PromptEmbeds:
|
||||||
text_embeds: torch.FloatTensor
|
text_embeds: torch.Tensor
|
||||||
pooled_embeds: Union[torch.FloatTensor, None]
|
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):
|
if isinstance(args, list) or isinstance(args, tuple):
|
||||||
# xl
|
# xl
|
||||||
self.text_embeds = args[0]
|
self.text_embeds = args[0]
|
||||||
@@ -29,30 +73,536 @@ class PromptEmbeds:
|
|||||||
self.text_embeds = args
|
self.text_embeds = args
|
||||||
self.pooled_embeds = None
|
self.pooled_embeds = None
|
||||||
|
|
||||||
def to(self, **kwargs):
|
def to(self, *args, **kwargs):
|
||||||
self.text_embeds = self.text_embeds.to(**kwargs)
|
self.text_embeds = self.text_embeds.to(*args, **kwargs)
|
||||||
if self.pooled_embeds is not None:
|
if self.pooled_embeds is not None:
|
||||||
self.pooled_embeds = self.pooled_embeds.to(**kwargs)
|
self.pooled_embeds = self.pooled_embeds.to(*args, **kwargs)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
from transformers import CLIPTextModel, CLIPTokenizer, CLIPTextModelWithProjection
|
||||||
|
|
||||||
|
# if is type checking
|
||||||
|
if typing.TYPE_CHECKING:
|
||||||
|
from diffusers import \
|
||||||
|
StableDiffusionPipeline, \
|
||||||
|
AutoencoderKL, \
|
||||||
|
UNet2DConditionModel
|
||||||
|
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||||
|
|
||||||
|
|
||||||
class StableDiffusion:
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
vae,
|
device,
|
||||||
tokenizer,
|
model_config: ModelConfig,
|
||||||
text_encoder,
|
dtype='fp16',
|
||||||
unet,
|
custom_pipeline=None,
|
||||||
noise_scheduler,
|
|
||||||
is_xl=False
|
|
||||||
):
|
):
|
||||||
# text encoder has a list of 2 for xl
|
self.custom_pipeline = custom_pipeline
|
||||||
self.vae = vae
|
self.device = device
|
||||||
self.tokenizer = tokenizer
|
self.dtype = dtype
|
||||||
self.text_encoder = text_encoder
|
self.torch_dtype = get_torch_dtype(dtype)
|
||||||
self.unet = unet
|
self.device_torch = torch.device(self.device)
|
||||||
self.noise_scheduler = noise_scheduler
|
self.model_config = model_config
|
||||||
self.is_xl = is_xl
|
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:
|
def encode_prompt(self, prompt, num_images_per_prompt=1) -> PromptEmbeds:
|
||||||
prompt = prompt
|
prompt = prompt
|
||||||
@@ -75,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):
|
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
|
# todo see what logit scale is
|
||||||
if self.is_xl:
|
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
|
# Convert the UNet model
|
||||||
update_sd("model.diffusion_model.", self.unet.state_dict())
|
update_sd("model.diffusion_model.", self.unet.state_dict())
|
||||||
|
|
||||||
@@ -96,19 +717,27 @@ class StableDiffusion:
|
|||||||
text_enc2_dict = convert_text_encoder_2_state_dict_to_sdxl(self.text_encoder[1].state_dict(), logit_scale)
|
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)
|
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
|
# Convert the VAE
|
||||||
|
if self.vae is not None:
|
||||||
vae_dict = model_util.convert_vae_state_dict(self.vae.state_dict())
|
vae_dict = model_util.convert_vae_state_dict(self.vae.state_dict())
|
||||||
update_sd("first_stage_model.", vae_dict)
|
update_sd("first_stage_model.", vae_dict)
|
||||||
|
|
||||||
# Put together new checkpoint
|
# prepare metadata
|
||||||
key_count = len(state_dict.keys())
|
meta = get_meta_for_safetensors(meta)
|
||||||
new_ckpt = {"state_dict": state_dict}
|
# make sure parent folder exists
|
||||||
|
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
||||||
if model_util.is_safetensors(output_file):
|
save_file(state_dict, output_file, metadata=meta)
|
||||||
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")
|
|
||||||
|
|||||||
316
toolkit/train_pipelines.py
Normal file
316
toolkit/train_pipelines.py
Normal file
@@ -0,0 +1,316 @@
|
|||||||
|
from typing import Optional, Tuple, Callable, Dict, Any, Union, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput
|
||||||
|
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||||
|
|
||||||
|
from toolkit.lora_special import LoRASpecialNetwork
|
||||||
|
from toolkit.pipelines import CustomStableDiffusionXLPipeline
|
||||||
|
|
||||||
|
|
||||||
|
class TransferStableDiffusionXLPipeline(CustomStableDiffusionXLPipeline):
|
||||||
|
def transfer_diffuse(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]] = None,
|
||||||
|
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
height: Optional[int] = None,
|
||||||
|
width: Optional[int] = None,
|
||||||
|
num_inference_steps: int = 50,
|
||||||
|
denoising_end: Optional[float] = None,
|
||||||
|
guidance_scale: float = 5.0,
|
||||||
|
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||||
|
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_images_per_prompt: Optional[int] = 1,
|
||||||
|
eta: float = 0.0,
|
||||||
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
|
latents: Optional[torch.FloatTensor] = None,
|
||||||
|
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
output_type: Optional[str] = "pil",
|
||||||
|
return_dict: bool = True,
|
||||||
|
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||||
|
callback_steps: int = 1,
|
||||||
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
guidance_rescale: float = 0.0,
|
||||||
|
original_size: Optional[Tuple[int, int]] = None,
|
||||||
|
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||||
|
target_size: Optional[Tuple[int, int]] = None,
|
||||||
|
target_unet: Optional[torch.nn.Module] = None,
|
||||||
|
pre_condition_callback = None,
|
||||||
|
each_step_callback = None,
|
||||||
|
network: Optional[LoRASpecialNetwork] = None,
|
||||||
|
):
|
||||||
|
r"""
|
||||||
|
Function invoked when calling the pipeline for generation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||||
|
instead.
|
||||||
|
prompt_2 (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||||
|
used in both text-encoders
|
||||||
|
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||||
|
The height in pixels of the generated image.
|
||||||
|
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||||
|
The width in pixels of the generated image.
|
||||||
|
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||||
|
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||||
|
expense of slower inference.
|
||||||
|
denoising_end (`float`, *optional*):
|
||||||
|
When specified, determines the fraction (between 0.0 and 1.0) of the total denoising process to be
|
||||||
|
completed before it is intentionally prematurely terminated. As a result, the returned sample will
|
||||||
|
still retain a substantial amount of noise as determined by the discrete timesteps selected by the
|
||||||
|
scheduler. The denoising_end parameter should ideally be utilized when this pipeline forms a part of a
|
||||||
|
"Mixture of Denoisers" multi-pipeline setup, as elaborated in [**Refining the Image
|
||||||
|
Output**](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl#refining-the-image-output)
|
||||||
|
guidance_scale (`float`, *optional*, defaults to 7.5):
|
||||||
|
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||||
|
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||||
|
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||||
|
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||||
|
usually at the expense of lower image quality.
|
||||||
|
negative_prompt (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||||
|
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||||
|
less than `1`).
|
||||||
|
negative_prompt_2 (`str` or `List[str]`, *optional*):
|
||||||
|
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
|
||||||
|
`text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders
|
||||||
|
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||||
|
The number of images to generate per prompt.
|
||||||
|
eta (`float`, *optional*, defaults to 0.0):
|
||||||
|
Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to
|
||||||
|
[`schedulers.DDIMScheduler`], will be ignored for others.
|
||||||
|
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||||
|
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||||
|
to make generation deterministic.
|
||||||
|
latents (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||||
|
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||||
|
tensor will ge generated by sampling using the supplied random `generator`.
|
||||||
|
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||||
|
provided, text embeddings will be generated from `prompt` input argument.
|
||||||
|
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||||
|
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||||
|
argument.
|
||||||
|
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||||
|
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||||
|
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||||
|
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||||
|
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
|
||||||
|
input argument.
|
||||||
|
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||||
|
The output format of the generate image. Choose between
|
||||||
|
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||||
|
return_dict (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead
|
||||||
|
of a plain tuple.
|
||||||
|
callback (`Callable`, *optional*):
|
||||||
|
A function that will be called every `callback_steps` steps during inference. The function will be
|
||||||
|
called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.
|
||||||
|
callback_steps (`int`, *optional*, defaults to 1):
|
||||||
|
The frequency at which the `callback` function will be called. If not specified, the callback will be
|
||||||
|
called at every step.
|
||||||
|
cross_attention_kwargs (`dict`, *optional*):
|
||||||
|
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||||
|
`self.processor` in
|
||||||
|
[diffusers.cross_attention](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/cross_attention.py).
|
||||||
|
guidance_rescale (`float`, *optional*, defaults to 0.7):
|
||||||
|
Guidance rescale factor proposed by [Common Diffusion Noise Schedules and Sample Steps are
|
||||||
|
Flawed](https://arxiv.org/pdf/2305.08891.pdf) `guidance_scale` is defined as `φ` in equation 16. of
|
||||||
|
[Common Diffusion Noise Schedules and Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf).
|
||||||
|
Guidance rescale factor should fix overexposure when using zero terminal SNR.
|
||||||
|
original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||||
|
If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled.
|
||||||
|
`original_size` defaults to `(width, height)` if not specified. Part of SDXL's micro-conditioning as
|
||||||
|
explained in section 2.2 of
|
||||||
|
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
|
||||||
|
`crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position
|
||||||
|
`crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting
|
||||||
|
`crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of
|
||||||
|
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||||
|
For most cases, `target_size` should be set to the desired height and width of the generated image. If
|
||||||
|
not specified it will default to `(width, height)`. Part of SDXL's micro-conditioning as explained in
|
||||||
|
section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] or `tuple`:
|
||||||
|
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] if `return_dict` is True, otherwise a
|
||||||
|
`tuple`. When returning a tuple, the first element is a list with the generated images.
|
||||||
|
"""
|
||||||
|
# 0. Default height and width to unet
|
||||||
|
height = height or self.default_sample_size * self.vae_scale_factor
|
||||||
|
width = width or self.default_sample_size * self.vae_scale_factor
|
||||||
|
|
||||||
|
original_size = original_size or (height, width)
|
||||||
|
target_size = target_size or (height, width)
|
||||||
|
|
||||||
|
# 1. Check inputs. Raise error if not correct
|
||||||
|
self.check_inputs(
|
||||||
|
prompt,
|
||||||
|
prompt_2,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
callback_steps,
|
||||||
|
negative_prompt,
|
||||||
|
negative_prompt_2,
|
||||||
|
prompt_embeds,
|
||||||
|
negative_prompt_embeds,
|
||||||
|
pooled_prompt_embeds,
|
||||||
|
negative_pooled_prompt_embeds,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Define call parameters
|
||||||
|
if prompt is not None and isinstance(prompt, str):
|
||||||
|
batch_size = 1
|
||||||
|
elif prompt is not None and isinstance(prompt, list):
|
||||||
|
batch_size = len(prompt)
|
||||||
|
else:
|
||||||
|
batch_size = prompt_embeds.shape[0]
|
||||||
|
|
||||||
|
device = self._execution_device
|
||||||
|
|
||||||
|
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||||
|
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||||
|
# corresponds to doing no classifier free guidance.
|
||||||
|
do_classifier_free_guidance = guidance_scale > 1.0
|
||||||
|
|
||||||
|
# 3. Encode input prompt
|
||||||
|
text_encoder_lora_scale = (
|
||||||
|
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
|
||||||
|
)
|
||||||
|
(
|
||||||
|
prompt_embeds,
|
||||||
|
negative_prompt_embeds,
|
||||||
|
pooled_prompt_embeds,
|
||||||
|
negative_pooled_prompt_embeds,
|
||||||
|
) = self.encode_prompt(
|
||||||
|
prompt=prompt,
|
||||||
|
prompt_2=prompt_2,
|
||||||
|
device=device,
|
||||||
|
num_images_per_prompt=num_images_per_prompt,
|
||||||
|
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||||
|
negative_prompt=negative_prompt,
|
||||||
|
negative_prompt_2=negative_prompt_2,
|
||||||
|
prompt_embeds=prompt_embeds,
|
||||||
|
negative_prompt_embeds=negative_prompt_embeds,
|
||||||
|
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||||
|
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||||
|
lora_scale=text_encoder_lora_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Prepare timesteps
|
||||||
|
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||||
|
|
||||||
|
timesteps = self.scheduler.timesteps
|
||||||
|
|
||||||
|
# 5. Prepare latent variables
|
||||||
|
num_channels_latents = self.unet.config.in_channels
|
||||||
|
latents = self.prepare_latents(
|
||||||
|
batch_size * num_images_per_prompt,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
prompt_embeds.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||||
|
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||||
|
|
||||||
|
# 7. Prepare added time ids & embeddings
|
||||||
|
add_text_embeds = pooled_prompt_embeds
|
||||||
|
add_time_ids = self._get_add_time_ids(
|
||||||
|
original_size, crops_coords_top_left, target_size, dtype=prompt_embeds.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||||
|
add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0)
|
||||||
|
add_time_ids = torch.cat([add_time_ids, add_time_ids], dim=0)
|
||||||
|
|
||||||
|
prompt_embeds = prompt_embeds.to(device)
|
||||||
|
add_text_embeds = add_text_embeds.to(device)
|
||||||
|
add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1)
|
||||||
|
|
||||||
|
# 8. Denoising loop
|
||||||
|
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||||
|
|
||||||
|
# 7.1 Apply denoising_end
|
||||||
|
if denoising_end is not None and type(denoising_end) == float and denoising_end > 0 and denoising_end < 1:
|
||||||
|
discrete_timestep_cutoff = int(
|
||||||
|
round(
|
||||||
|
self.scheduler.config.num_train_timesteps
|
||||||
|
- (denoising_end * self.scheduler.config.num_train_timesteps)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
num_inference_steps = len(list(filter(lambda ts: ts >= discrete_timestep_cutoff, timesteps)))
|
||||||
|
timesteps = timesteps[:num_inference_steps]
|
||||||
|
|
||||||
|
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||||
|
for i, t in enumerate(timesteps):
|
||||||
|
# expand the latents if we are doing classifier free guidance
|
||||||
|
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||||
|
|
||||||
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids}
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input,
|
||||||
|
t,
|
||||||
|
encoder_hidden_states=prompt_embeds,
|
||||||
|
cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
added_cond_kwargs=added_cond_kwargs,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
conditioned_noise_pred, conditioned_latent_model_input = pre_condition_callback(
|
||||||
|
noise_pred.clone().detach(),
|
||||||
|
latent_model_input.clone().detach(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# start grad
|
||||||
|
with torch.enable_grad():
|
||||||
|
with network:
|
||||||
|
assert network.is_active
|
||||||
|
noise_train_pred = target_unet(
|
||||||
|
conditioned_latent_model_input,
|
||||||
|
t,
|
||||||
|
encoder_hidden_states=prompt_embeds,
|
||||||
|
cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
added_cond_kwargs=added_cond_kwargs,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
each_step_callback(conditioned_noise_pred, noise_train_pred)
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||||
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||||
|
|
||||||
|
if do_classifier_free_guidance and 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)
|
||||||
|
|
||||||
|
# compute the previous noisy sample x_t -> x_t-1
|
||||||
|
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||||
|
|
||||||
|
# call the callback, if provided
|
||||||
|
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||||
|
progress_bar.update()
|
||||||
|
if callback is not None and i % callback_steps == 0:
|
||||||
|
callback(i, t, latents)
|
||||||
|
|
||||||
@@ -1,8 +1,13 @@
|
|||||||
import argparse
|
import argparse
|
||||||
|
import hashlib
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
import sys
|
||||||
|
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||||
|
|
||||||
|
sys.path.append(SD_SCRIPTS_ROOT)
|
||||||
|
|
||||||
from diffusers import (
|
from diffusers import (
|
||||||
StableDiffusionPipeline,
|
StableDiffusionPipeline,
|
||||||
@@ -29,13 +34,16 @@ SCHEDLER_SCHEDULE = "scaled_linear"
|
|||||||
|
|
||||||
|
|
||||||
def get_torch_dtype(dtype_str):
|
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":
|
if dtype_str == "float" or dtype_str == "fp32" or dtype_str == "single" or dtype_str == "float32":
|
||||||
return torch.float
|
return torch.float
|
||||||
if dtype_str == "fp16" or dtype_str == "half" or dtype_str == "float16":
|
if dtype_str == "fp16" or dtype_str == "half" or dtype_str == "float16":
|
||||||
return torch.float16
|
return torch.float16
|
||||||
if dtype_str == "bf16" or dtype_str == "bfloat16":
|
if dtype_str == "bf16" or dtype_str == "bfloat16":
|
||||||
return torch.bfloat16
|
return torch.bfloat16
|
||||||
return None
|
return dtype_str
|
||||||
|
|
||||||
|
|
||||||
def replace_filewords_prompt(prompt, args: argparse.Namespace):
|
def replace_filewords_prompt(prompt, args: argparse.Namespace):
|
||||||
@@ -399,3 +407,29 @@ def concat_prompt_embeddings(
|
|||||||
[unconditional.pooled_embeds, conditional.pooled_embeds]
|
[unconditional.pooled_embeds, conditional.pooled_embeds]
|
||||||
).repeat_interleave(n_imgs, dim=0)
|
).repeat_interleave(n_imgs, dim=0)
|
||||||
return PromptEmbeds([text_embeds, pooled_embeds])
|
return PromptEmbeds([text_embeds, pooled_embeds])
|
||||||
|
|
||||||
|
|
||||||
|
def addnet_hash_safetensors(b):
|
||||||
|
"""New model hash used by sd-webui-additional-networks for .safetensors format files"""
|
||||||
|
hash_sha256 = hashlib.sha256()
|
||||||
|
blksize = 1024 * 1024
|
||||||
|
|
||||||
|
b.seek(0)
|
||||||
|
header = b.read(8)
|
||||||
|
n = int.from_bytes(header, "little")
|
||||||
|
|
||||||
|
offset = n + 8
|
||||||
|
b.seek(offset)
|
||||||
|
for chunk in iter(lambda: b.read(blksize), b""):
|
||||||
|
hash_sha256.update(chunk)
|
||||||
|
|
||||||
|
return hash_sha256.hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def addnet_hash_legacy(b):
|
||||||
|
"""Old model hash used by sd-webui-additional-networks for .safetensors format files"""
|
||||||
|
m = hashlib.sha256()
|
||||||
|
|
||||||
|
b.seek(0x100000)
|
||||||
|
m.update(b.read(0x10000))
|
||||||
|
return m.hexdigest()[0:8]
|
||||||
|
|||||||
Reference in New Issue
Block a user