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/_PUT_YOUR_CONFIGS_HERE).txt
|
||||
/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
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
# or source venv/Scripts/activate on windows
|
||||
# .\venv\Scripts\activate on windows
|
||||
# windows install pytorch first with
|
||||
# pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu117
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
@@ -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
|
||||
here so far.
|
||||
|
||||
---
|
||||
|
||||
### Batch Image Generation
|
||||
|
||||
A image generator that can take frompts from a config file or form a txt file and generate them to a
|
||||
folder. I mainly needed this for an SDXL test I am doing but added some polish to it so it can be used
|
||||
for generat batch image generation.
|
||||
It all runs off a config file, which you can find an example of in `config/examples/generate.example.yaml`.
|
||||
Mere info is in the comments in the example
|
||||
|
||||
---
|
||||
|
||||
### LoRA (lierla), LoCON (LyCORIS) extractor
|
||||
|
||||
It is based on the extractor in the [LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS) tool, but adding some QOL features
|
||||
@@ -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.
|
||||
|
||||
---
|
||||
|
||||
### 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
|
||||
|
||||
<a target="_blank" href="https://colab.research.google.com/github/ostris/ai-toolkit/blob/main/notebooks/SliderTraining.ipynb">
|
||||
<img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/>
|
||||
</a>
|
||||
|
||||
This is how I train most of the recent sliders I have on Civitai, you can check them out in my [Civitai profile](https://civitai.com/user/Ostris/models).
|
||||
It is based off the work by [p1atdev/LECO](https://github.com/p1atdev/LECO) and [rohitgandikota/erasing](https://github.com/rohitgandikota/erasing)
|
||||
But has been heavily modified to create sliders rather than erasing concepts. I have a lot more plans on this, but it is
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -108,14 +168,53 @@ Just went in and out. It is much worse on smaller faces than shown here.
|
||||
|
||||
## TODO
|
||||
- [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
|
||||
- [ ] Make Textual inversion network trainer (network that spits out TI embeddings)
|
||||
|
||||
---
|
||||
|
||||
## 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
|
||||
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:
|
||||
# the name will be used to create a folder in the output folder
|
||||
# it will also replace any [name] token in the rest of this config
|
||||
name: pet_slider_v1
|
||||
name: detail_slider_v1
|
||||
# folder will be created with name above in folder below
|
||||
# it can be relative to the project root or absolute
|
||||
training_folder: "output/LoRA"
|
||||
@@ -24,7 +24,7 @@ config:
|
||||
type: "lierla"
|
||||
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
|
||||
rank: 8
|
||||
alpha: 1.0 # just leave it
|
||||
alpha: 4 # Do about half of rank
|
||||
|
||||
# training config
|
||||
train:
|
||||
@@ -33,7 +33,9 @@ config:
|
||||
# how many steps to train. More is not always better. I rarely go over 1000
|
||||
steps: 500
|
||||
# I have had good results with 4e-4 to 1e-4 at 500 steps
|
||||
lr: 1e-4
|
||||
lr: 2e-4
|
||||
# enables gradient checkpoint, saves vram, leave it on
|
||||
gradient_checkpointing: true
|
||||
# train the unet. I recommend leaving this true
|
||||
train_unet: true
|
||||
# train the text encoder. I don't recommend this unless you have a special use case
|
||||
@@ -41,6 +43,7 @@ config:
|
||||
# not the description of it (text encoder)
|
||||
train_text_encoder: false
|
||||
|
||||
|
||||
# just leave unless you know what you are doing
|
||||
# also supports "dadaptation" but set lr to 1 if you use that,
|
||||
# but it learns too fast and I don't recommend it
|
||||
@@ -51,11 +54,13 @@ config:
|
||||
# while training. Just leave it
|
||||
max_denoising_steps: 40
|
||||
# works great at 1. I do 1 even with my 4090.
|
||||
# higher may not work right with newer single batch stacking code anyway
|
||||
batch_size: 1
|
||||
# bf16 works best if your GPU supports it (modern)
|
||||
dtype: bf16 # fp32, bf16, fp16
|
||||
# if you have it, use it. It is faster and better
|
||||
xformers: true
|
||||
# torch 2.0 doesnt need xformers anymore, only use if you have lower version
|
||||
# xformers: true
|
||||
# I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX
|
||||
# although, the way we train sliders is comparative, so it probably won't work anyway
|
||||
noise_offset: 0.0
|
||||
@@ -66,11 +71,17 @@ config:
|
||||
name_or_path: "runwayml/stable-diffusion-v1-5"
|
||||
is_v2: false # for v2 models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
# has some issues with the dual text encoder and the way we train sliders
|
||||
# it works bit weights need to probably be higher to see it.
|
||||
is_xl: false # for SDXL models
|
||||
|
||||
# saving config
|
||||
save:
|
||||
dtype: float16 # precision to save. I recommend float16
|
||||
save_every: 50 # save every this many steps
|
||||
# this will remove step counts more than this number
|
||||
# allows you to save more often in case of a crash without filling up your drive
|
||||
max_step_saves_to_keep: 2
|
||||
|
||||
# sampling config
|
||||
sample:
|
||||
@@ -88,21 +99,22 @@ config:
|
||||
# --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive
|
||||
# slide are good tests. will inherit sample.network_multiplier if not set
|
||||
# --n [string] # negative prompt, will inherit sample.neg if not set
|
||||
|
||||
# Only 75 tokens allowed currently
|
||||
prompts: # our example is an animal slider, neg: dog, pos: cat
|
||||
- "a golden retriever --m -5"
|
||||
- "a golden retriever --m -3"
|
||||
- "a golden retriever --m 3"
|
||||
- "a golden retriever --m 5"
|
||||
- "calico cat --m -5"
|
||||
- "calico cat --m -3"
|
||||
- "calico cat --m 3"
|
||||
- "calico cat --m 5"
|
||||
- "an elephant --m -5"
|
||||
- "an elephant --m -3"
|
||||
- "an elephant --m 3"
|
||||
- "an elephant --m 5"
|
||||
# I like to do a wide positive and negative spread so I can see a good range and stop
|
||||
# early if the network is braking down
|
||||
prompts:
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3"
|
||||
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5"
|
||||
- "a golden retriever sitting on a leather couch, --m -5"
|
||||
- "a golden retriever sitting on a leather couch --m -3"
|
||||
- "a golden retriever sitting on a leather couch --m 3"
|
||||
- "a golden retriever sitting on a leather couch --m 5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3"
|
||||
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5"
|
||||
# negative prompt used on all prompts above as default if they don't have one
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
|
||||
# seed for sampling. 42 is the answer for everything
|
||||
@@ -131,11 +143,16 @@ config:
|
||||
# resolutions to train on. [ width, height ]. This is less important for sliders
|
||||
# as we are not teaching the model anything it doesn't already know
|
||||
# but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1
|
||||
# and [ 1024, 1024 ] for sd_xl
|
||||
# you can do as many as you want here
|
||||
resolutions:
|
||||
- [ 512, 512 ]
|
||||
# - [ 512, 768 ]
|
||||
# - [ 768, 768 ]
|
||||
# slider training uses 4 combined steps for a single round. This will do it in one gradient
|
||||
# step. It is highly optimized and shouldn't take anymore vram than doing without it,
|
||||
# since we break down batches for gradient accumulation now. so just leave it on.
|
||||
batch_full_slide: true
|
||||
# These are the concepts to train on. You can do as many as you want here,
|
||||
# but they can conflict outweigh each other. Other than experimenting, I recommend
|
||||
# just doing one for good results
|
||||
@@ -146,7 +163,9 @@ config:
|
||||
# a keyword necessarily but what the model understands the concept to represent.
|
||||
# "person" will affect men, women, children, etc but will not affect cats, dogs, etc
|
||||
# it is the models base general understanding of the concept and everything it represents
|
||||
- target_class: "animal"
|
||||
# you can leave it blank to affect everything. In this example, we are adjusting
|
||||
# detail, so we will leave it blank to affect everything
|
||||
- target_class: ""
|
||||
# positive is the prompt for the positive side of the slider.
|
||||
# It is the concept that will be excited and amplified in the model when we slide the slider
|
||||
# to the positive side and forgotten / inverted when we slide
|
||||
@@ -154,33 +173,44 @@ config:
|
||||
# the prompt. You want it to be the extreme of what you want to train on. For example,
|
||||
# if you want to train on fat people, you would use "an extremely fat, morbidly obese person"
|
||||
# as the prompt. Not just "fat person"
|
||||
positive: "cat"
|
||||
# max 75 tokens for now
|
||||
positive: "high detail, 8k, intricate, detailed, high resolution, high res, high quality"
|
||||
# negative is the prompt for the negative side of the slider and works the same as positive
|
||||
# it does not necessarily work the same as a negative prompt when generating images
|
||||
negative: "dog"
|
||||
# these need to be polar opposites.
|
||||
# max 76 tokens for now
|
||||
negative: "blurry, boring, fuzzy, low detail, low resolution, low res, low quality"
|
||||
# the loss for this target is multiplied by this number.
|
||||
# if you are doing more than one target it may be good to set less important ones
|
||||
# to a lower number like 0.1 so they dont outweigh the primary target
|
||||
# to a lower number like 0.1 so they don't outweigh the primary target
|
||||
weight: 1.0
|
||||
|
||||
# anchors are prompts that wer try to hold on to while training the slider
|
||||
# you want these to generate an image very similar to the target_class
|
||||
# without directly overlapping it. For example, if you are training on a person smiling,
|
||||
# you would use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not
|
||||
anchors:
|
||||
# only positive prompt for now
|
||||
- prompt: "a woman"
|
||||
neg_prompt: "animal"
|
||||
# the multiplier applied to the LoRA when this is run.
|
||||
# higher will give it more weight but also help keep the lora from collapsing
|
||||
multiplier: 8.0
|
||||
- prompt: "a man"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
- prompt: "a person"
|
||||
neg_prompt: "animal"
|
||||
multiplier: 8.0
|
||||
|
||||
# anchors are prompts that we will try to hold on to while training the slider
|
||||
# these are NOT necessary and can prevent the slider from converging if not done right
|
||||
# leave them off if you are having issues, but they can help lock the network
|
||||
# on certain concepts to help prevent catastrophic forgetting
|
||||
# you want these to generate an image that is not your target_class, but close to it
|
||||
# is fine as long as it does not directly overlap it.
|
||||
# For example, if you are training on a person smiling,
|
||||
# you could use "a person with a face mask" as an anchor. It is a person, the image is the same
|
||||
# regardless if they are smiling or not, however, the closer the concept is to the target_class
|
||||
# the less the multiplier needs to be. Keep multipliers less than 1.0 for anchors usually
|
||||
# for close concepts, you want to be closer to 0.1 or 0.2
|
||||
# these will slow down training. I am leaving them off for the demo
|
||||
|
||||
# anchors:
|
||||
# - prompt: "a woman"
|
||||
# neg_prompt: "animal"
|
||||
# # the multiplier applied to the LoRA when this is run.
|
||||
# # higher will give it more weight but also help keep the lora from collapsing
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a man"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
# - prompt: "a person"
|
||||
# neg_prompt: "animal"
|
||||
# multiplier: 1.0
|
||||
|
||||
# You can put any information you want here, and it will be saved in the model.
|
||||
# The below is an example, but you can put your grocery list in it if you want.
|
||||
|
||||
129
extensions/example/ExampleMergeModels.py
Normal file
129
extensions/example/ExampleMergeModels.py
Normal file
@@ -0,0 +1,129 @@
|
||||
import torch
|
||||
import gc
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from tqdm import tqdm
|
||||
|
||||
# Type check imports. Prevents circular imports
|
||||
if TYPE_CHECKING:
|
||||
from jobs import ExtensionJob
|
||||
|
||||
|
||||
# extend standard config classes to add weight
|
||||
class ModelInputConfig(ModelConfig):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.weight = kwargs.get('weight', 1.0)
|
||||
# overwrite default dtype unless user specifies otherwise
|
||||
# float 32 will give up better precision on the merging functions
|
||||
self.dtype: str = kwargs.get('dtype', 'float32')
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
# this is our main class process
|
||||
class ExampleMergeModels(BaseExtensionProcess):
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job: 'ExtensionJob',
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
# this is the setup process, do not do process intensive stuff here, just variable setup and
|
||||
# checking requirements. This is called before the run() function
|
||||
# no loading models or anything like that, it is just for setting up the process
|
||||
# all of your process intensive stuff should be done in the run() function
|
||||
# config will have everything from the process item in the config file
|
||||
|
||||
# convince methods exist on BaseProcess to get config values
|
||||
# if required is set to true and the value is not found it will throw an error
|
||||
# you can pass a default value to get_conf() as well if it was not in the config file
|
||||
# as well as a type to cast the value to
|
||||
self.save_path = self.get_conf('save_path', required=True)
|
||||
self.save_dtype = self.get_conf('save_dtype', default='float16', as_type=get_torch_dtype)
|
||||
self.device = self.get_conf('device', default='cpu', as_type=torch.device)
|
||||
|
||||
# build models to merge list
|
||||
models_to_merge = self.get_conf('models_to_merge', required=True, as_type=list)
|
||||
# build list of ModelInputConfig objects. I find it is a good idea to make a class for each config
|
||||
# this way you can add methods to it and it is easier to read and code. There are a lot of
|
||||
# inbuilt config classes located in toolkit.config_modules as well
|
||||
self.models_to_merge = [ModelInputConfig(**model) for model in models_to_merge]
|
||||
# setup is complete. Don't load anything else here, just setup variables and stuff
|
||||
|
||||
# this is the entire run process be sure to call super().run() first
|
||||
def run(self):
|
||||
# always call first
|
||||
super().run()
|
||||
print(f"Running process: {self.__class__.__name__}")
|
||||
|
||||
# let's adjust our weights first to normalize them so the total is 1.0
|
||||
total_weight = sum([model.weight for model in self.models_to_merge])
|
||||
weight_adjust = 1.0 / total_weight
|
||||
for model in self.models_to_merge:
|
||||
model.weight *= weight_adjust
|
||||
|
||||
output_model: StableDiffusion = None
|
||||
# let's do the merge, it is a good idea to use tqdm to show progress
|
||||
for model_config in tqdm(self.models_to_merge, desc="Merging models"):
|
||||
# setup model class with our helper class
|
||||
sd_model = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=model_config,
|
||||
dtype="float32"
|
||||
)
|
||||
# load the model
|
||||
sd_model.load_model()
|
||||
|
||||
# adjust the weight of the text encoder
|
||||
if isinstance(sd_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for text_encoder in sd_model.text_encoder:
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
else:
|
||||
# normal model
|
||||
for key, value in sd_model.text_encoder.state_dict().items():
|
||||
value *= model_config.weight
|
||||
# adjust the weights of the unet
|
||||
for key, value in sd_model.unet.state_dict().items():
|
||||
value *= model_config.weight
|
||||
|
||||
if output_model is None:
|
||||
# use this one as the base
|
||||
output_model = sd_model
|
||||
else:
|
||||
# merge the models
|
||||
# text encoder
|
||||
if isinstance(output_model.text_encoder, list):
|
||||
# sdxl model
|
||||
for i, text_encoder in enumerate(output_model.text_encoder):
|
||||
for key, value in text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder[i].state_dict()[key]
|
||||
else:
|
||||
# normal model
|
||||
for key, value in output_model.text_encoder.state_dict().items():
|
||||
value += sd_model.text_encoder.state_dict()[key]
|
||||
# unet
|
||||
for key, value in output_model.unet.state_dict().items():
|
||||
value += sd_model.unet.state_dict()[key]
|
||||
|
||||
# remove the model to free memory
|
||||
del sd_model
|
||||
flush()
|
||||
|
||||
# merge loop is done, let's save the model
|
||||
print(f"Saving merged model to {self.save_path}")
|
||||
output_model.save(self.save_path, meta=self.meta, save_dtype=self.save_dtype)
|
||||
print(f"Saved merged model to {self.save_path}")
|
||||
# do cleanup here
|
||||
del output_model
|
||||
flush()
|
||||
25
extensions/example/__init__.py
Normal file
25
extensions/example/__init__.py
Normal file
@@ -0,0 +1,25 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# We make a subclass of Extension
|
||||
class ExampleMergeExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "example_merge_extension"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Example Merge Extension"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .ExampleMergeModels import ExampleMergeModels
|
||||
return ExampleMergeModels
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ExampleMergeExtension
|
||||
]
|
||||
48
extensions/example/config/config.example.yaml
Normal file
48
extensions/example/config/config.example.yaml
Normal file
@@ -0,0 +1,48 @@
|
||||
---
|
||||
# Always include at least one example config file to show how to use your extension.
|
||||
# use plenty of comments so users know how to use it and what everything does
|
||||
|
||||
# all extensions will use this job name
|
||||
job: extension
|
||||
config:
|
||||
name: 'my_awesome_merge'
|
||||
process:
|
||||
# Put your example processes here. This will be passed
|
||||
# to your extension process in the config argument.
|
||||
# the type MUST match your extension uid
|
||||
- type: "example_merge_extension"
|
||||
# save path for the merged model
|
||||
save_path: "output/merge/[name].safetensors"
|
||||
# save type
|
||||
dtype: fp16
|
||||
# device to run it on
|
||||
device: cuda:0
|
||||
# input models can only be SD1.x and SD2.x models for this example (currently)
|
||||
models_to_merge:
|
||||
# weights are relative, total weights will be normalized
|
||||
# for example. If you have 2 models with weight 1.0, they will
|
||||
# both be weighted 0.5. If you have 1 model with weight 1.0 and
|
||||
# another with weight 2.0, the first will be weighted 1/3 and the
|
||||
# second will be weighted 2/3
|
||||
- name_or_path: "input/model1.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model2.safetensors"
|
||||
weight: 1.0
|
||||
- name_or_path: "input/model3.safetensors"
|
||||
weight: 0.3
|
||||
- name_or_path: "input/model4.safetensors"
|
||||
weight: 1.0
|
||||
|
||||
|
||||
# you can put any information you want here, and it will be saved in the model
|
||||
# the below is an example. I recommend doing trigger words at a minimum
|
||||
# in the metadata. The software will include this plus some other information
|
||||
meta:
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your model
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
@@ -0,0 +1,218 @@
|
||||
import copy
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from typing import Optional, Union, List
|
||||
from torch.utils.data import ConcatDataset, DataLoader
|
||||
from toolkit.data_loader import PairedImageDataset
|
||||
from toolkit.prompt_utils import concat_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class DatasetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.pair_folder: str = kwargs.get('pair_folder', None)
|
||||
self.network_weight: float = kwargs.get('network_weight', 1.0)
|
||||
self.target_class: str = kwargs.get('target_class', '')
|
||||
self.size: int = kwargs.get('size', 512)
|
||||
|
||||
|
||||
class ReferenceSliderConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
|
||||
self.datasets: List[DatasetConfig] = [DatasetConfig(**d) for d in kwargs.get('datasets', [])]
|
||||
|
||||
|
||||
class ImageReferenceSliderTrainerProcess(BaseSDTrainProcess):
|
||||
sd: StableDiffusion
|
||||
data_loader: DataLoader = None
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
self.prompt_txt_list = None
|
||||
self.step_num = 0
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.slider_config = ReferenceSliderConfig(**self.get_conf('slider', {}))
|
||||
|
||||
def load_datasets(self):
|
||||
if self.data_loader is None:
|
||||
print(f"Loading datasets")
|
||||
datasets = []
|
||||
for dataset in self.slider_config.datasets:
|
||||
print(f" - Dataset: {dataset.pair_folder}")
|
||||
config = {
|
||||
'path': dataset.pair_folder,
|
||||
'size': dataset.size,
|
||||
'default_prompt': dataset.target_class,
|
||||
'network_weight': dataset.network_weight,
|
||||
}
|
||||
image_dataset = PairedImageDataset(config)
|
||||
datasets.append(image_dataset)
|
||||
|
||||
concatenated_dataset = ConcatDataset(datasets)
|
||||
self.data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=self.train_config.batch_size,
|
||||
shuffle=True,
|
||||
num_workers=2
|
||||
)
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
self.load_datasets()
|
||||
|
||||
pass
|
||||
|
||||
def hook_train_loop(self, batch):
|
||||
do_mirror_loss = 'mirror' in self.slider_config.additional_losses
|
||||
|
||||
with torch.no_grad():
|
||||
imgs, prompts, base_network_weight = batch
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
imgs: torch.Tensor = imgs.to(self.device_torch, dtype=dtype)
|
||||
# split batched images in half so left is negative and right is positive
|
||||
negative_images, positive_images = torch.chunk(imgs, 2, dim=3)
|
||||
|
||||
positive_latents = self.sd.encode_images(positive_images)
|
||||
negative_latents = self.sd.encode_images(negative_images)
|
||||
|
||||
height = positive_images.shape[2]
|
||||
width = positive_images.shape[3]
|
||||
batch_size = positive_images.shape[0]
|
||||
|
||||
if self.train_config.gradient_checkpointing:
|
||||
# may get disabled elsewhere
|
||||
self.sd.unet.enable_gradient_checkpointing()
|
||||
|
||||
noise_scheduler = self.sd.noise_scheduler
|
||||
optimizer = self.optimizer
|
||||
lr_scheduler = self.lr_scheduler
|
||||
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
|
||||
timesteps = torch.randint(0, self.train_config.max_denoising_steps, (1,), device=self.device_torch)
|
||||
timesteps = timesteps.long()
|
||||
|
||||
# get noise
|
||||
noise_positive = self.sd.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
batch_size=batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
if do_mirror_loss:
|
||||
# mirror the noise
|
||||
# torch shape is [batch, channels, height, width]
|
||||
noise_negative = torch.flip(noise_positive.clone(), dims=[3])
|
||||
else:
|
||||
noise_negative = noise_positive.clone()
|
||||
|
||||
# Add noise to the latents according to the noise magnitude at each timestep
|
||||
# (this is the forward diffusion process)
|
||||
noisy_positive_latents = noise_scheduler.add_noise(positive_latents, noise_positive, timesteps)
|
||||
noisy_negative_latents = noise_scheduler.add_noise(negative_latents, noise_negative, timesteps)
|
||||
|
||||
noisy_latents = torch.cat([noisy_positive_latents, noisy_negative_latents], dim=0)
|
||||
noise = torch.cat([noise_positive, noise_negative], dim=0)
|
||||
timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
network_multiplier = [base_network_weight * 1.0, base_network_weight * -1.0]
|
||||
|
||||
flush()
|
||||
|
||||
loss_float = None
|
||||
loss_slide_float = None
|
||||
loss_mirror_float = None
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
noisy_latents.requires_grad = False
|
||||
|
||||
# if training text encoder enable grads, else do context of no grad
|
||||
with torch.set_grad_enabled(self.train_config.train_text_encoder):
|
||||
# text encoding
|
||||
embedding_list = []
|
||||
# embed the prompts
|
||||
for prompt in prompts:
|
||||
embedding = self.sd.encode_prompt(prompt).to(self.device_torch, dtype=dtype)
|
||||
embedding_list.append(embedding)
|
||||
conditional_embeds = concat_prompt_embeds(embedding_list)
|
||||
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
|
||||
self.network.multiplier = network_multiplier
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents,
|
||||
conditional_embeddings=conditional_embeds,
|
||||
timestep=timesteps,
|
||||
)
|
||||
|
||||
if self.sd.prediction_type == 'v_prediction':
|
||||
# v-parameterization training
|
||||
target = noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
|
||||
else:
|
||||
target = noise
|
||||
|
||||
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
# todo add snr gamma here
|
||||
|
||||
loss = loss.mean()
|
||||
loss_slide_float = loss.item()
|
||||
|
||||
if do_mirror_loss:
|
||||
noise_pred_pos, noise_pred_neg = torch.chunk(noise_pred, 2, dim=0)
|
||||
# mirror the negative
|
||||
noise_pred_neg = torch.flip(noise_pred_neg.clone(), dims=[3])
|
||||
loss_mirror = torch.nn.functional.mse_loss(noise_pred_pos.float(), noise_pred_neg.float(),
|
||||
reduction="none")
|
||||
loss_mirror = loss_mirror.mean([1, 2, 3])
|
||||
loss_mirror = loss_mirror.mean()
|
||||
loss_mirror_float = loss_mirror.item()
|
||||
loss += loss_mirror
|
||||
|
||||
loss_float = loss.item()
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
|
||||
flush()
|
||||
|
||||
# apply gradients
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss_float},
|
||||
)
|
||||
|
||||
if do_mirror_loss:
|
||||
loss_dict['l/s'] = loss_slide_float
|
||||
loss_dict['l/m'] = loss_mirror_float
|
||||
return loss_dict
|
||||
# end hook_train_loop
|
||||
@@ -0,0 +1,25 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# We make a subclass of Extension
|
||||
class ImageReferenceSliderTrainer(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "image_reference_slider_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Image Reference Slider Trainer"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .ImageReferenceSliderTrainerProcess import ImageReferenceSliderTrainerProcess
|
||||
return ImageReferenceSliderTrainerProcess
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ImageReferenceSliderTrainer
|
||||
]
|
||||
@@ -0,0 +1,107 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: example_name
|
||||
process:
|
||||
- type: 'image_reference_slider_trainer'
|
||||
training_folder: "/mnt/Train/out/LoRA"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "/home/jaret/Dev/.tensorboard"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 8
|
||||
linear_alpha: 8
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 5000
|
||||
lr: 1e-4
|
||||
train_unet: true
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: true
|
||||
optimizer: "adamw"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 1
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
skip_first_sample: true
|
||||
noise_offset: 0.0
|
||||
model:
|
||||
name_or_path: "/path/to/model.safetensors"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 1000 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # only affects step counts
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of a woman with red hair taking a selfie --m -3"
|
||||
- "photo of a woman with red hair taking a selfie --m -1"
|
||||
- "photo of a woman with red hair taking a selfie --m 1"
|
||||
- "photo of a woman with red hair taking a selfie --m 3"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m -3"
|
||||
- "close up photo of a man smiling at the camera, in a tank top--m -1"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m 1"
|
||||
- "close up photo of a man smiling at the camera, in a tank top --m 3"
|
||||
- "photo of a blonde woman smiling, barista --m -3"
|
||||
- "photo of a blonde woman smiling, barista --m -1"
|
||||
- "photo of a blonde woman smiling, barista --m 1"
|
||||
- "photo of a blonde woman smiling, barista --m 3"
|
||||
- "photo of a Christina Hendricks --m -1"
|
||||
- "photo of a Christina Hendricks --m -1"
|
||||
- "photo of a Christina Hendricks --m 1"
|
||||
- "photo of a Christina Hendricks --m 3"
|
||||
- "photo of a Christina Ricci --m -3"
|
||||
- "photo of a Christina Ricci --m -1"
|
||||
- "photo of a Christina Ricci --m 1"
|
||||
- "photo of a Christina Ricci --m 3"
|
||||
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
|
||||
seed: 42
|
||||
walk_seed: false
|
||||
guidance_scale: 7
|
||||
sample_steps: 20
|
||||
network_multiplier: 1.0
|
||||
|
||||
logging:
|
||||
log_every: 10 # log every this many steps
|
||||
use_wandb: false # not supported yet
|
||||
verbose: false
|
||||
|
||||
slider:
|
||||
datasets:
|
||||
- pair_folder: "/path/to/folder/side/by/side/images"
|
||||
network_weight: 2.0
|
||||
target_class: "" # only used as default if caption txt are not present
|
||||
size: 512
|
||||
- pair_folder: "/path/to/folder/side/by/side/images"
|
||||
network_weight: 4.0
|
||||
target_class: "" # only used as default if caption txt are not present
|
||||
size: 512
|
||||
|
||||
|
||||
# you can put any information you want here, and it will be saved in the model
|
||||
# the below is an example. I recommend doing trigger words at a minimum
|
||||
# in the metadata. The software will include this plus some other information
|
||||
meta:
|
||||
name: "[name]" # [name] gets replaced with the name above
|
||||
description: A short description of your model
|
||||
trigger_words:
|
||||
- put
|
||||
- trigger
|
||||
- words
|
||||
- here
|
||||
version: '0.1'
|
||||
creator:
|
||||
name: Your Name
|
||||
email: your@email.com
|
||||
website: https://yourwebsite.com
|
||||
any: All meta data above is arbitrary, it can be whatever you want.
|
||||
2
info.py
2
info.py
@@ -3,6 +3,6 @@ from collections import OrderedDict
|
||||
v = OrderedDict()
|
||||
v["name"] = "ai-toolkit"
|
||||
v["repo"] = "https://github.com/ostris/ai-toolkit"
|
||||
v["version"] = "0.0.1"
|
||||
v["version"] = "0.0.4"
|
||||
|
||||
software_meta = v
|
||||
|
||||
@@ -60,7 +60,11 @@ class BaseJob:
|
||||
|
||||
# check if dict key is process type
|
||||
if process['type'] in process_dict:
|
||||
ProcessClass = getattr(module, process_dict[process['type']])
|
||||
if isinstance(process_dict[process['type']], str):
|
||||
ProcessClass = getattr(module, process_dict[process['type']])
|
||||
else:
|
||||
# it is the class
|
||||
ProcessClass = process_dict[process['type']]
|
||||
self.process.append(ProcessClass(i, self, process))
|
||||
else:
|
||||
raise ValueError(f'config file is invalid. Unknown process type: {process["type"]}')
|
||||
|
||||
21
jobs/ExtensionJob.py
Normal file
21
jobs/ExtensionJob.py
Normal file
@@ -0,0 +1,21 @@
|
||||
from collections import OrderedDict
|
||||
from jobs import BaseJob
|
||||
from toolkit.extension import get_all_extensions_process_dict
|
||||
|
||||
|
||||
class ExtensionJob(BaseJob):
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
super().__init__(config)
|
||||
self.device = self.get_conf('device', 'cpu')
|
||||
self.process_dict = get_all_extensions_process_dict()
|
||||
self.load_processes(self.process_dict)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
print("")
|
||||
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
32
jobs/GenerateJob.py
Normal file
32
jobs/GenerateJob.py
Normal file
@@ -0,0 +1,32 @@
|
||||
from jobs import BaseJob
|
||||
from collections import OrderedDict
|
||||
from typing import List
|
||||
from jobs.process import GenerateProcess
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
|
||||
import sys
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
process_dict = {
|
||||
'to_folder': 'GenerateProcess',
|
||||
}
|
||||
|
||||
|
||||
class GenerateJob(BaseJob):
|
||||
process: List[GenerateProcess]
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
super().__init__(config)
|
||||
self.device = self.get_conf('device', 'cpu')
|
||||
|
||||
# loads the processes from the config
|
||||
self.load_processes(process_dict)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("")
|
||||
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
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 = {
|
||||
'vae': 'TrainVAEProcess',
|
||||
'slider': 'TrainSliderProcess',
|
||||
'slider_old': 'TrainSliderProcessOld',
|
||||
'lora_hack': 'TrainLoRAHack',
|
||||
'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.log_dir = self.get_conf('log_dir', None)
|
||||
|
||||
self.writer = None
|
||||
self.setup_tensorboard()
|
||||
|
||||
# loads the processes from the config
|
||||
self.load_processes(process_dict)
|
||||
|
||||
def save_training_config(self):
|
||||
timestamp = datetime.now().strftime('%Y%m%d-%H%M%S')
|
||||
os.makedirs(self.training_folder, exist_ok=True)
|
||||
save_dif = os.path.join(self.training_folder, f'run_config_{timestamp}.yaml')
|
||||
with open(save_dif, 'w') as f:
|
||||
yaml.dump(self.raw_config, f)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -54,12 +48,3 @@ class TrainJob(BaseJob):
|
||||
|
||||
for process in self.process:
|
||||
process.run()
|
||||
|
||||
def setup_tensorboard(self):
|
||||
if self.log_dir:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
now = datetime.now()
|
||||
time_str = now.strftime('%Y%m%d-%H%M%S')
|
||||
summary_name = f"{self.name}_{time_str}"
|
||||
summary_dir = os.path.join(self.log_dir, summary_name)
|
||||
self.writer = SummaryWriter(summary_dir)
|
||||
|
||||
@@ -2,3 +2,6 @@ from .BaseJob import BaseJob
|
||||
from .ExtractJob import ExtractJob
|
||||
from .TrainJob import TrainJob
|
||||
from .MergeJob import MergeJob
|
||||
from .ModJob import ModJob
|
||||
from .GenerateJob import GenerateJob
|
||||
from .ExtensionJob import ExtensionJob
|
||||
|
||||
20
jobs/process/BaseExtensionProcess.py
Normal file
20
jobs/process/BaseExtensionProcess.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
|
||||
|
||||
class BaseExtensionProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job,
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -1,10 +1,9 @@
|
||||
import copy
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
|
||||
|
||||
class BaseProcess:
|
||||
class BaseProcess(object):
|
||||
meta: OrderedDict
|
||||
|
||||
def __init__(
|
||||
@@ -16,6 +15,8 @@ class BaseProcess:
|
||||
self.process_id = process_id
|
||||
self.job = job
|
||||
self.config = config
|
||||
self.raw_process_config = config
|
||||
self.name = self.get_conf('name', self.job.name)
|
||||
self.meta = copy.deepcopy(self.job.meta)
|
||||
print(json.dumps(self.config, indent=4))
|
||||
|
||||
|
||||
@@ -1,34 +1,26 @@
|
||||
import glob
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from typing import Union
|
||||
|
||||
import diffusers
|
||||
from safetensors import safe_open
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
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.optimizer import get_optimizer
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, DDPMScheduler
|
||||
from toolkit.scheduler import get_lr_scheduler
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from jobs.process import BaseTrainProcess
|
||||
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
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
|
||||
import gc
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from leco import train_util, model_util
|
||||
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \
|
||||
GenerateImageConfig
|
||||
|
||||
|
||||
def flush():
|
||||
@@ -36,212 +28,104 @@ def flush():
|
||||
gc.collect()
|
||||
|
||||
|
||||
UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも同じ。
|
||||
VAE_SCALE_FACTOR = 8 # 2 ** (len(vae.config.block_out_channels) - 1) = 8
|
||||
|
||||
|
||||
class BaseSDTrainProcess(BaseTrainProcess):
|
||||
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)
|
||||
self.custom_pipeline = custom_pipeline
|
||||
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.network_config = NetworkConfig(**self.get_conf('network', None))
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
||||
network_config = self.get_conf('network', None)
|
||||
if network_config is not None:
|
||||
self.network_config = NetworkConfig(**network_config)
|
||||
else:
|
||||
self.network_config = None
|
||||
self.train_config = TrainConfig(**self.get_conf('train', {}))
|
||||
self.model_config = ModelConfig(**self.get_conf('model', {}))
|
||||
self.save_config = SaveConfig(**self.get_conf('save', {}))
|
||||
self.sample_config = SampleConfig(**self.get_conf('sample', {}))
|
||||
self.first_sample_config = SampleConfig(
|
||||
**self.get_conf('first_sample', {})) if 'first_sample' in self.config else self.sample_config
|
||||
first_sample_config = self.get_conf('first_sample', None)
|
||||
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.optimizer = None
|
||||
self.lr_scheduler = None
|
||||
self.sd: 'StableDiffusion' = None
|
||||
self.data_loader: Union[DataLoader, None] = None
|
||||
|
||||
# sdxl stuff
|
||||
self.logit_scale = None
|
||||
self.ckppt_info = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.train_config.dtype,
|
||||
custom_pipeline=self.custom_pipeline,
|
||||
)
|
||||
|
||||
# added later
|
||||
# to hold network if there is one
|
||||
self.network = None
|
||||
|
||||
def sample(self, step=None, is_first=False):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
os.makedirs(sample_folder, exist_ok=True)
|
||||
|
||||
if self.network is not None:
|
||||
self.network.eval()
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
original_device_dict = {
|
||||
'vae': self.sd.vae.device,
|
||||
'unet': self.sd.unet.device,
|
||||
# 'tokenizer': self.sd.tokenizer.device,
|
||||
}
|
||||
|
||||
# handle sdxl text encoder
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder, i in zip(self.sd.text_encoder, range(len(self.sd.text_encoder))):
|
||||
original_device_dict[f'text_encoder_{i}'] = encoder.device
|
||||
encoder.to(self.device_torch)
|
||||
else:
|
||||
original_device_dict['text_encoder'] = self.sd.text_encoder.device
|
||||
self.sd.text_encoder.to(self.device_torch)
|
||||
|
||||
self.sd.vae.to(self.device_torch)
|
||||
self.sd.unet.to(self.device_torch)
|
||||
# self.sd.text_encoder.to(self.device_torch)
|
||||
# self.sd.tokenizer.to(self.device_torch)
|
||||
# TODO add clip skip
|
||||
if self.sd.is_xl:
|
||||
pipeline = StableDiffusionXLPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=self.sd.noise_scheduler,
|
||||
)
|
||||
else:
|
||||
pipeline = StableDiffusionPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder,
|
||||
tokenizer=self.sd.tokenizer,
|
||||
scheduler=self.sd.noise_scheduler,
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
)
|
||||
# disable progress bar
|
||||
pipeline.set_progress_bar_config(disable=True)
|
||||
gen_img_config_list = []
|
||||
|
||||
sample_config = self.first_sample_config if is_first else self.sample_config
|
||||
|
||||
start_seed = sample_config.seed
|
||||
start_multiplier = self.network.multiplier
|
||||
current_seed = start_seed
|
||||
for i in range(len(sample_config.prompts)):
|
||||
if sample_config.walk_seed:
|
||||
current_seed = start_seed + i
|
||||
|
||||
pipeline.to(self.device_torch)
|
||||
with self.network:
|
||||
with torch.no_grad():
|
||||
if self.network is not None:
|
||||
assert self.network.is_active
|
||||
if self.logging_config.verbose:
|
||||
print("network_state", {
|
||||
'is_active': self.network.is_active,
|
||||
'multiplier': self.network.multiplier,
|
||||
})
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
|
||||
for i in tqdm(range(len(sample_config.prompts)), desc=f"Generating Samples - step: {step}",
|
||||
leave=False):
|
||||
raw_prompt = sample_config.prompts[i]
|
||||
filename = f"[time]_{step_num}_[count].png"
|
||||
|
||||
neg = sample_config.neg
|
||||
multiplier = sample_config.network_multiplier
|
||||
p_split = raw_prompt.split('--')
|
||||
prompt = p_split[0].strip()
|
||||
height = sample_config.height
|
||||
width = sample_config.width
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
|
||||
if len(p_split) > 1:
|
||||
for split in p_split:
|
||||
flag = split[:1]
|
||||
content = split[1:].strip()
|
||||
if flag == 'n':
|
||||
neg = content
|
||||
elif flag == 'm':
|
||||
# multiplier
|
||||
multiplier = float(content)
|
||||
elif flag == 'w':
|
||||
# multiplier
|
||||
width = int(content)
|
||||
elif flag == 'h':
|
||||
# multiplier
|
||||
height = int(content)
|
||||
gen_img_config_list.append(GenerateImageConfig(
|
||||
prompt=sample_config.prompts[i], # it will autoparse the prompt
|
||||
width=sample_config.width,
|
||||
height=sample_config.height,
|
||||
negative_prompt=sample_config.neg,
|
||||
seed=current_seed,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
guidance_rescale=sample_config.guidance_rescale,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
network_multiplier=sample_config.network_multiplier,
|
||||
output_path=output_path,
|
||||
))
|
||||
|
||||
height = max(64, height - height % 8) # round to divisible by 8
|
||||
width = max(64, width - width % 8) # round to divisible by 8
|
||||
|
||||
if sample_config.walk_seed:
|
||||
current_seed += i
|
||||
|
||||
if self.network is not None:
|
||||
self.network.multiplier = multiplier
|
||||
torch.manual_seed(current_seed)
|
||||
torch.cuda.manual_seed(current_seed)
|
||||
|
||||
if self.sd.is_xl:
|
||||
img = pipeline(
|
||||
prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
negative_prompt=neg,
|
||||
).images[0]
|
||||
else:
|
||||
img = pipeline(
|
||||
prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_inference_steps=sample_config.sample_steps,
|
||||
guidance_scale=sample_config.guidance_scale,
|
||||
negative_prompt=neg,
|
||||
).images[0]
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
seconds_since_epoch = int(time.time())
|
||||
# zero-pad 2 digits
|
||||
i_str = str(i).zfill(2)
|
||||
filename = f"{seconds_since_epoch}{step_num}_{i_str}.png"
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
img.save(output_path)
|
||||
|
||||
# clear pipeline and cache to reduce vram usage
|
||||
del pipeline
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# restore training state
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
|
||||
self.sd.vae.to(original_device_dict['vae'])
|
||||
self.sd.unet.to(original_device_dict['unet'])
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder, i in zip(self.sd.text_encoder, range(len(self.sd.text_encoder))):
|
||||
encoder.to(original_device_dict[f'text_encoder_{i}'])
|
||||
else:
|
||||
self.sd.text_encoder.to(original_device_dict['text_encoder'])
|
||||
if self.network is not None:
|
||||
self.network.train()
|
||||
self.network.multiplier = start_multiplier
|
||||
# self.sd.tokenizer.to(original_device_dict['tokenizer'])
|
||||
# send to be generated
|
||||
self.sd.generate_images(gen_img_config_list)
|
||||
|
||||
def update_training_metadata(self):
|
||||
dict = OrderedDict({
|
||||
o_dict = OrderedDict({
|
||||
"training_info": self.get_training_info()
|
||||
})
|
||||
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:
|
||||
dict['ss_base_model_version'] = 'sdxl_1.0'
|
||||
elif self.model_config.is_xl:
|
||||
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):
|
||||
info = OrderedDict({
|
||||
@@ -299,6 +183,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
)
|
||||
|
||||
self.print(f"Saved to {file_path}")
|
||||
self.clean_up_saves()
|
||||
|
||||
# Called before the model is loaded
|
||||
def hook_before_model_load(self):
|
||||
@@ -312,153 +197,10 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
def hook_before_train_loop(self):
|
||||
pass
|
||||
|
||||
def get_latent_noise(
|
||||
self,
|
||||
height=None,
|
||||
width=None,
|
||||
pixel_height=None,
|
||||
pixel_width=None,
|
||||
):
|
||||
if height is None and pixel_height is None:
|
||||
raise ValueError("height or pixel_height must be specified")
|
||||
raise ValueError("height or pixel_height must be specified")
|
||||
if width is None and pixel_width is None:
|
||||
raise ValueError("width or pixel_width must be specified")
|
||||
if height is None:
|
||||
height = pixel_height // VAE_SCALE_FACTOR
|
||||
if width is None:
|
||||
width = pixel_width // VAE_SCALE_FACTOR
|
||||
|
||||
noise = torch.randn(
|
||||
(
|
||||
self.train_config.batch_size,
|
||||
UNET_IN_CHANNELS,
|
||||
height,
|
||||
width,
|
||||
),
|
||||
device="cpu",
|
||||
)
|
||||
noise = apply_noise_offset(noise, self.train_config.noise_offset)
|
||||
return noise
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch=None):
|
||||
# return loss
|
||||
return 0.0
|
||||
|
||||
def get_time_ids_from_latents(self, latents):
|
||||
bs, ch, h, w = list(latents.shape)
|
||||
|
||||
height = h * VAE_SCALE_FACTOR
|
||||
width = w * VAE_SCALE_FACTOR
|
||||
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
if self.sd.is_xl:
|
||||
prompt_ids = train_util.get_add_time_ids(
|
||||
height,
|
||||
width,
|
||||
dynamic_crops=False, # look into this
|
||||
dtype=dtype,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
return train_util.concat_embeddings(
|
||||
prompt_ids, prompt_ids, bs
|
||||
)
|
||||
else:
|
||||
return None
|
||||
|
||||
def predict_noise(
|
||||
self,
|
||||
latents: torch.FloatTensor,
|
||||
text_embeddings: PromptEmbeds,
|
||||
timestep: int,
|
||||
guidance_scale=7.5,
|
||||
guidance_rescale=0.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):
|
||||
# get latest saved step
|
||||
if os.path.exists(self.save_root):
|
||||
@@ -486,84 +228,56 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
print("load_weights not implemented for non-network models")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
# run base process run
|
||||
BaseTrainProcess.run(self)
|
||||
### HOOK ###
|
||||
self.hook_before_model_load()
|
||||
# run base sd process run
|
||||
self.sd.load_model()
|
||||
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
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:
|
||||
vae.set_use_memory_efficient_attention_xformers(True)
|
||||
unet.enable_xformers_memory_efficient_attention()
|
||||
if self.train_config.gradient_checkpointing:
|
||||
unet.enable_gradient_checkpointing()
|
||||
# if isinstance(text_encoder, list):
|
||||
# for te in text_encoder:
|
||||
# te.enable_gradient_checkpointing()
|
||||
# else:
|
||||
# text_encoder.enable_gradient_checkpointing()
|
||||
|
||||
unet.to(self.device_torch, dtype=dtype)
|
||||
unet.requires_grad_(False)
|
||||
unet.eval()
|
||||
vae = vae.to(torch.device('cpu'), dtype=dtype)
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
|
||||
if self.network_config is not None:
|
||||
conv = self.network_config.conv if self.network_config.conv is not None and self.network_config.conv > 0 else None
|
||||
self.network = LoRASpecialNetwork(
|
||||
text_encoder=text_encoder,
|
||||
unet=unet,
|
||||
lora_dim=self.network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=self.network_config.alpha,
|
||||
alpha=self.network_config.linear_alpha,
|
||||
train_unet=self.train_config.train_unet,
|
||||
train_text_encoder=self.train_config.train_text_encoder,
|
||||
conv_lora_dim=conv,
|
||||
conv_alpha=self.network_config.alpha if conv is not None else None,
|
||||
conv_lora_dim=self.network_config.conv,
|
||||
conv_alpha=self.network_config.conv_alpha,
|
||||
)
|
||||
|
||||
self.network.force_to(self.device_torch, dtype=dtype)
|
||||
# give network to sd so it can use it
|
||||
self.sd.network = self.network
|
||||
|
||||
self.network.apply_to(
|
||||
text_encoder,
|
||||
@@ -580,6 +294,9 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
default_lr=self.train_config.lr
|
||||
)
|
||||
|
||||
if self.train_config.gradient_checkpointing:
|
||||
self.network.enable_gradient_checkpointing()
|
||||
|
||||
latest_save_path = self.get_latest_save_path()
|
||||
if latest_save_path is not None:
|
||||
self.print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
|
||||
@@ -588,14 +305,19 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
|
||||
|
||||
else:
|
||||
params = []
|
||||
# assume dreambooth/finetune
|
||||
if self.train_config.train_text_encoder:
|
||||
text_encoder.requires_grad_(True)
|
||||
text_encoder.train()
|
||||
params += text_encoder.parameters()
|
||||
if self.sd.is_xl:
|
||||
for te in text_encoder:
|
||||
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:
|
||||
unet.requires_grad_(True)
|
||||
unet.train()
|
||||
@@ -609,11 +331,11 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
optimizer_params=self.train_config.optimizer_params)
|
||||
self.optimizer = optimizer
|
||||
|
||||
lr_scheduler = train_util.get_lr_scheduler(
|
||||
lr_scheduler = get_lr_scheduler(
|
||||
self.train_config.lr_scheduler,
|
||||
optimizer,
|
||||
max_iterations=self.train_config.steps,
|
||||
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
|
||||
@@ -621,28 +343,51 @@ class BaseSDTrainProcess(BaseTrainProcess):
|
||||
### HOOK ###
|
||||
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
|
||||
if self.train_config.skip_first_sample:
|
||||
self.print("Skipping first sample due to config setting")
|
||||
else:
|
||||
self.print("Generating baseline samples before training")
|
||||
self.sample(0, is_first=True)
|
||||
self.sample(0)
|
||||
|
||||
self.progress_bar = tqdm(
|
||||
total=self.train_config.steps,
|
||||
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
|
||||
for step in range(self.step_num, self.train_config.steps):
|
||||
# todo handle dataloader here maybe, not sure
|
||||
if dataloader is not None:
|
||||
try:
|
||||
batch = next(dataloader_iterator)
|
||||
except StopIteration:
|
||||
# hit the end of an epoch, reset
|
||||
# todo, should we do something else here? like blow up balloons?
|
||||
dataloader_iterator = iter(dataloader)
|
||||
batch = next(dataloader_iterator)
|
||||
else:
|
||||
batch = None
|
||||
|
||||
### HOOK ###
|
||||
loss_dict = self.hook_train_loop()
|
||||
loss_dict = self.hook_train_loop(batch)
|
||||
flush()
|
||||
|
||||
if self.train_config.optimizer.lower().startswith('dadaptation'):
|
||||
if self.train_config.optimizer.lower().startswith('dadaptation') or \
|
||||
self.train_config.optimizer.lower().startswith('prodigy'):
|
||||
learning_rate = (
|
||||
optimizer.param_groups[0]["d"] *
|
||||
optimizer.param_groups[0]["lr"]
|
||||
|
||||
@@ -1,14 +1,24 @@
|
||||
from datetime import datetime
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
import yaml
|
||||
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jobs import TrainJob, BaseJob, ExtensionJob
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class BaseTrainProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
writer: 'SummaryWriter'
|
||||
job: Union['TrainJob', 'BaseJob', 'ExtensionJob']
|
||||
progress_bar: 'tqdm' = None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -18,11 +28,14 @@ class BaseTrainProcess(BaseProcess):
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.progress_bar = None
|
||||
self.writer = self.job.writer
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
|
||||
self.save_root = os.path.join(self.training_folder, self.job.name)
|
||||
self.writer = None
|
||||
self.training_folder = self.get_conf('training_folder', self.job.training_folder if hasattr(self.job, 'training_folder') else None)
|
||||
self.save_root = os.path.join(self.training_folder, self.name)
|
||||
self.step = 0
|
||||
self.first_step = 0
|
||||
self.log_dir = self.get_conf('log_dir', self.job.log_dir if hasattr(self.job, 'log_dir') else None)
|
||||
self.setup_tensorboard()
|
||||
self.save_training_config()
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -37,3 +50,19 @@ class BaseTrainProcess(BaseProcess):
|
||||
self.progress_bar.update()
|
||||
else:
|
||||
print(*args)
|
||||
|
||||
def setup_tensorboard(self):
|
||||
if self.log_dir:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
now = datetime.now()
|
||||
time_str = now.strftime('%Y%m%d-%H%M%S')
|
||||
summary_name = f"{self.name}_{time_str}"
|
||||
summary_dir = os.path.join(self.log_dir, summary_name)
|
||||
self.writer = SummaryWriter(summary_dir)
|
||||
|
||||
def save_training_config(self):
|
||||
timestamp = datetime.now().strftime('%Y%m%d-%H%M%S')
|
||||
os.makedirs(self.save_root, exist_ok=True)
|
||||
save_dif = os.path.join(self.save_root, f'process_config_{timestamp}.yaml')
|
||||
with open(save_dif, 'w') as f:
|
||||
yaml.dump(self.raw_process_config, f)
|
||||
|
||||
102
jobs/process/GenerateProcess.py
Normal file
102
jobs/process/GenerateProcess.py
Normal file
@@ -0,0 +1,102 @@
|
||||
import gc
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef, List
|
||||
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig
|
||||
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_model_hash_to_meta, \
|
||||
add_base_model_info_to_meta
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
prompts: List[str]
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.width = kwargs.get('width', 512)
|
||||
self.height = kwargs.get('height', 512)
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.prompt_2 = kwargs.get('prompt_2', None)
|
||||
self.neg_2 = kwargs.get('neg_2', None)
|
||||
self.prompts = kwargs.get('prompts', None)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.prompt_file = kwargs.get('prompt_file', False)
|
||||
if self.prompts is None:
|
||||
raise ValueError("Prompts must be set")
|
||||
if isinstance(self.prompts, str):
|
||||
if os.path.exists(self.prompts):
|
||||
with open(self.prompts, 'r', encoding='utf-8') as f:
|
||||
self.prompts = f.read().splitlines()
|
||||
self.prompts = [p.strip() for p in self.prompts if len(p.strip()) > 0]
|
||||
else:
|
||||
raise ValueError("Prompts file does not exist, put in list if you want to use a list of prompts")
|
||||
|
||||
|
||||
class GenerateProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
progress_bar: ForwardRef('tqdm') = None
|
||||
sd: StableDiffusion
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job,
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.output_folder = self.get_conf('output_folder', required=True)
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.model_config.dtype,
|
||||
)
|
||||
print(f"Using device {self.device}")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for prompt in self.generate_config.prompts:
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=self.generate_config.width,
|
||||
height=self.generate_config.height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs)
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
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
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
if self.hack_config.type == 'suppression':
|
||||
return self.supress_loop()
|
||||
else:
|
||||
|
||||
@@ -1,22 +1,14 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import glob
|
||||
import os
|
||||
from typing import Optional
|
||||
from collections import OrderedDict
|
||||
import random
|
||||
from typing import Optional, List
|
||||
|
||||
from safetensors.torch import load_file, save_file
|
||||
from safetensors.torch import save_file, load_file
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.layers import ReductionKernel
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||
|
||||
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
|
||||
@@ -38,12 +30,10 @@ class RescaleConfig:
|
||||
):
|
||||
self.from_resolution = kwargs.get('from_resolution', 512)
|
||||
self.scale = kwargs.get('scale', 0.5)
|
||||
self.prompt_file = kwargs.get('prompt_file', None)
|
||||
self.prompt_tensors = kwargs.get('prompt_tensors', None)
|
||||
self.latent_tensor_dir = kwargs.get('latent_tensor_dir', None)
|
||||
self.num_latent_tensors = kwargs.get('num_latent_tensors', 1000)
|
||||
self.to_resolution = kwargs.get('to_resolution', int(self.from_resolution * self.scale))
|
||||
|
||||
if self.prompt_file is None:
|
||||
raise ValueError("prompt_file is required")
|
||||
self.prompt_dropout = kwargs.get('prompt_dropout', 0.1)
|
||||
|
||||
|
||||
class PromptEmbedsCache:
|
||||
@@ -61,12 +51,12 @@ class PromptEmbedsCache:
|
||||
|
||||
class TrainSDRescaleProcess(BaseSDTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
# pass our custom pipeline to super so it sets it up
|
||||
super().__init__(process_id, job, config)
|
||||
self.step_num = 0
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.prompt_cache = PromptEmbedsCache()
|
||||
self.rescale_config = RescaleConfig(**self.get_conf('rescale', required=True))
|
||||
self.reduce_size_fn = ReductionKernel(
|
||||
in_channels=4,
|
||||
@@ -74,202 +64,211 @@ class TrainSDRescaleProcess(BaseSDTrainProcess):
|
||||
dtype=get_torch_dtype(self.train_config.dtype),
|
||||
device=self.device_torch,
|
||||
)
|
||||
self.prompt_txt_list = []
|
||||
|
||||
self.latent_paths: List[str] = []
|
||||
self.empty_embedding: PromptEmbeds = None
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def get_latent_tensors(self):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
num_to_generate = 0
|
||||
# check if dir exists
|
||||
if not os.path.exists(self.rescale_config.latent_tensor_dir):
|
||||
os.makedirs(self.rescale_config.latent_tensor_dir)
|
||||
num_to_generate = self.rescale_config.num_latent_tensors
|
||||
else:
|
||||
# find existing
|
||||
current_tensor_list = glob.glob(os.path.join(self.rescale_config.latent_tensor_dir, "*.safetensors"))
|
||||
num_to_generate = self.rescale_config.num_latent_tensors - len(current_tensor_list)
|
||||
self.latent_paths = current_tensor_list
|
||||
|
||||
if num_to_generate > 0:
|
||||
print(f"Generating {num_to_generate}/{self.rescale_config.num_latent_tensors} latent tensors")
|
||||
|
||||
# unload other model
|
||||
self.sd.unet.to('cpu')
|
||||
|
||||
# load aux network
|
||||
self.sd_parent = StableDiffusion(
|
||||
self.device_torch,
|
||||
model_config=self.model_config,
|
||||
dtype=self.train_config.dtype,
|
||||
)
|
||||
self.sd_parent.load_model()
|
||||
self.sd_parent.unet.to(self.device_torch, dtype=dtype)
|
||||
# we dont need text encoder for this
|
||||
|
||||
del self.sd_parent.text_encoder
|
||||
del self.sd_parent.tokenizer
|
||||
|
||||
self.sd_parent.unet.eval()
|
||||
self.sd_parent.unet.requires_grad_(False)
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
self.empty_embedding, # unconditional (negative prompt)
|
||||
self.empty_embedding, # conditional (positive prompt)
|
||||
self.train_config.batch_size,
|
||||
)
|
||||
torch.set_default_device(self.device_torch)
|
||||
|
||||
for i in tqdm(range(num_to_generate)):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
# get a random seed
|
||||
seed = torch.randint(0, 2 ** 32, (1,)).item()
|
||||
# zero pad seed string to max length
|
||||
seed_string = str(seed).zfill(10)
|
||||
# set seed
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
# # ger a random number of steps
|
||||
timesteps_to = self.train_config.max_denoising_steps
|
||||
|
||||
# set the scheduler to the number of steps
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
timesteps_to, device=self.device_torch
|
||||
)
|
||||
|
||||
noise = self.sd.get_latent_noise(
|
||||
pixel_height=self.rescale_config.from_resolution,
|
||||
pixel_width=self.rescale_config.from_resolution,
|
||||
batch_size=self.train_config.batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get latents
|
||||
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
||||
latents = latents.to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get random guidance scale from 1.0 to 10.0 (CFG)
|
||||
guidance_scale = torch.rand(1).item() * 9.0 + 1.0
|
||||
|
||||
# do a timestep of 1
|
||||
timestep = 1
|
||||
|
||||
noise_pred_target = self.sd_parent.predict_noise(
|
||||
latents,
|
||||
text_embeddings=text_embeddings,
|
||||
timestep=timestep,
|
||||
guidance_scale=guidance_scale
|
||||
)
|
||||
|
||||
# build state dict
|
||||
state_dict = OrderedDict()
|
||||
state_dict['noise_pred_target'] = noise_pred_target.to('cpu', dtype=torch.float16)
|
||||
state_dict['latents'] = latents.to('cpu', dtype=torch.float16)
|
||||
state_dict['guidance_scale'] = torch.tensor(guidance_scale).to('cpu', dtype=torch.float16)
|
||||
state_dict['timestep'] = torch.tensor(timestep).to('cpu', dtype=torch.float16)
|
||||
state_dict['timesteps_to'] = torch.tensor(timesteps_to).to('cpu', dtype=torch.float16)
|
||||
state_dict['seed'] = torch.tensor(seed).to('cpu', dtype=torch.float32) # must be float 32 to prevent overflow
|
||||
|
||||
file_name = f"{seed_string}_{i}.safetensors"
|
||||
file_path = os.path.join(self.rescale_config.latent_tensor_dir, file_name)
|
||||
save_file(state_dict, file_path)
|
||||
self.latent_paths.append(file_path)
|
||||
|
||||
print("Removing parent model")
|
||||
# delete parent
|
||||
del self.sd_parent
|
||||
flush()
|
||||
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
self.sd.unet.to(self.device_torch, dtype=dtype)
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.print(f"Loading prompt file from {self.rescale_config.prompt_file}")
|
||||
# encode our empty prompt
|
||||
self.empty_embedding = self.sd.encode_prompt("")
|
||||
self.empty_embedding = self.empty_embedding.to(self.device_torch,
|
||||
dtype=get_torch_dtype(self.train_config.dtype))
|
||||
|
||||
# read line by line from file
|
||||
with open(self.rescale_config.prompt_file, 'r') as f:
|
||||
self.prompt_txt_list = f.readlines()
|
||||
# clean empty lines
|
||||
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
|
||||
|
||||
self.print(f"Loaded {len(self.prompt_txt_list)} prompts. Encoding them..")
|
||||
|
||||
cache = PromptEmbedsCache()
|
||||
|
||||
# get encoded latents for our prompts
|
||||
with torch.no_grad():
|
||||
if self.rescale_config.prompt_tensors is not None:
|
||||
# check to see if it exists
|
||||
if os.path.exists(self.rescale_config.prompt_tensors):
|
||||
# load it.
|
||||
self.print(f"Loading prompt tensors from {self.rescale_config.prompt_tensors}")
|
||||
prompt_tensors = load_file(self.rescale_config.prompt_tensors, device='cpu')
|
||||
# add them to the cache
|
||||
for prompt_txt, prompt_tensor in prompt_tensors.items():
|
||||
if prompt_txt.startswith("te:"):
|
||||
prompt = prompt_txt[3:]
|
||||
# text_embeds
|
||||
text_embeds = prompt_tensor
|
||||
pooled_embeds = None
|
||||
# find pool embeds
|
||||
if f"pe:{prompt}" in prompt_tensors:
|
||||
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
|
||||
|
||||
# make it
|
||||
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
|
||||
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
|
||||
|
||||
if len(cache.prompts) == 0:
|
||||
print("Prompt tensors not found. Encoding prompts..")
|
||||
neutral = ""
|
||||
# encode neutral
|
||||
cache[neutral] = self.sd.encode_prompt(neutral)
|
||||
for prompt in tqdm(self.prompt_txt_list, desc="Encoding prompts", leave=False):
|
||||
# build the cache
|
||||
if cache[prompt] is None:
|
||||
cache[prompt] = self.sd.encode_prompt(prompt).to(device="cpu", dtype=torch.float32)
|
||||
|
||||
if self.rescale_config.prompt_tensors:
|
||||
print(f"Saving prompt tensors to {self.rescale_config.prompt_tensors}")
|
||||
state_dict = {}
|
||||
for prompt_txt, prompt_embeds in cache.prompts.items():
|
||||
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to("cpu", dtype=get_torch_dtype('fp16'))
|
||||
if prompt_embeds.pooled_embeds is not None:
|
||||
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to("cpu", dtype=get_torch_dtype('fp16'))
|
||||
save_file(state_dict, self.rescale_config.prompt_tensors)
|
||||
|
||||
self.print("Encoding complete.")
|
||||
|
||||
# move to cpu to save vram
|
||||
# We don't need text encoder anymore, but keep it on cpu for sampling
|
||||
# if text encoder is list
|
||||
# Move train model encoder to cpu
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for encoder in self.sd.text_encoder:
|
||||
encoder.to("cpu")
|
||||
encoder.to('cpu')
|
||||
encoder.eval()
|
||||
encoder.requires_grad_(False)
|
||||
else:
|
||||
self.sd.text_encoder.to("cpu")
|
||||
self.prompt_cache = cache
|
||||
self.sd.text_encoder.to('cpu')
|
||||
self.sd.text_encoder.eval()
|
||||
self.sd.text_encoder.requires_grad_(False)
|
||||
|
||||
# self.sd.unet.to('cpu')
|
||||
flush()
|
||||
|
||||
self.get_latent_tensors()
|
||||
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# get random encoded prompt from cache
|
||||
prompt_txt = self.prompt_txt_list[
|
||||
torch.randint(0, len(self.prompt_txt_list), (1,)).item()
|
||||
]
|
||||
prompt = self.prompt_cache[prompt_txt].to(device=self.device_torch, dtype=dtype)
|
||||
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()
|
||||
|
||||
def get_noise_pred(p, n, gs, cts, dn):
|
||||
return self.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,
|
||||
)
|
||||
# train it
|
||||
# Begin gradient accumulation
|
||||
self.sd.unet.train()
|
||||
self.sd.unet.requires_grad_(True)
|
||||
self.sd.unet.to(self.device_torch, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
self.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()
|
||||
# pick random latent tensor
|
||||
latent_path = random.choice(self.latent_paths)
|
||||
latent_tensor = load_file(latent_path)
|
||||
|
||||
# get noise
|
||||
noise = self.get_latent_noise(
|
||||
pixel_height=self.rescale_config.from_resolution,
|
||||
pixel_width=self.rescale_config.from_resolution,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
noise_pred_target = (latent_tensor['noise_pred_target']).to(self.device_torch, dtype=dtype)
|
||||
latents = (latent_tensor['latents']).to(self.device_torch, dtype=dtype)
|
||||
guidance_scale = (latent_tensor['guidance_scale']).item()
|
||||
timestep = int((latent_tensor['timestep']).item())
|
||||
timesteps_to = int((latent_tensor['timesteps_to']).item())
|
||||
# seed = int((latent_tensor['seed']).item())
|
||||
|
||||
# get latents
|
||||
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
||||
latents = latents.to(self.device_torch, dtype=dtype)
|
||||
#
|
||||
# # predict without network
|
||||
# assert self.network.is_active is False
|
||||
# denoised_latents = self.diffuse_some_steps(
|
||||
# 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
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
self.empty_embedding, # unconditional (negative prompt)
|
||||
self.empty_embedding, # conditional (positive prompt)
|
||||
self.train_config.batch_size,
|
||||
)
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
timesteps_to, device=self.device_torch
|
||||
)
|
||||
|
||||
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
|
||||
to_denoised_latents = self.reduce_size_fn(denoised_latents)
|
||||
# get the reduced latents
|
||||
# reduced_pred = self.reduce_size_fn(noise_pred_target.detach())
|
||||
denoised_target = self.reduce_size_fn(denoised_target.detach())
|
||||
reduced_latents = self.reduce_size_fn(latents.detach())
|
||||
|
||||
# start gradient
|
||||
optimizer.zero_grad()
|
||||
self.network.multiplier = 1.0
|
||||
with self.network:
|
||||
assert self.network.is_active is True
|
||||
to_prediction = get_noise_pred(
|
||||
prompt, neutral, 1, current_timestep, to_denoised_latents
|
||||
).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_target.requires_grad = False
|
||||
self.optimizer.zero_grad()
|
||||
noise_pred_train = self.sd.predict_noise(
|
||||
reduced_latents,
|
||||
text_embeddings=text_embeddings,
|
||||
timestep=timestep,
|
||||
guidance_scale=guidance_scale
|
||||
)
|
||||
|
||||
denoised_pred = self.sd.noise_scheduler.step(noise_pred_train, timestep, reduced_latents).prev_sample
|
||||
loss = loss_function(denoised_pred, denoised_target)
|
||||
loss_float = loss.item()
|
||||
|
||||
loss = loss.to(self.device_torch)
|
||||
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
del (
|
||||
reduced_from_prediction,
|
||||
from_prediction,
|
||||
to_denoised_latents,
|
||||
to_prediction,
|
||||
latents,
|
||||
)
|
||||
flush()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss_float},
|
||||
)
|
||||
|
||||
@@ -1,30 +1,31 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import time
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Optional, Union
|
||||
|
||||
from safetensors.torch import save_file, load_file
|
||||
import torch.utils.checkpoint as cp
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.layers import CheckpointGradients
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
from toolkit.stable_diffusion_model import PromptEmbeds
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
from toolkit.prompt_utils import \
|
||||
EncodedPromptPair, ACTION_TYPES_SLIDER, \
|
||||
EncodedAnchor, concat_prompt_pairs, \
|
||||
concat_anchors, PromptEmbedsCache, encode_prompts_to_cache, build_prompt_pair_batch_from_cache, split_anchors, \
|
||||
split_prompt_pairs
|
||||
|
||||
import torch
|
||||
from leco import train_util, model_util
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
|
||||
|
||||
|
||||
class ACTION_TYPES_SLIDER:
|
||||
ERASE_NEGATIVE = 0
|
||||
ENHANCE_NEGATIVE = 1
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
@@ -32,58 +33,10 @@ def flush():
|
||||
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):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
self.prompt_txt_list = None
|
||||
self.step_num = 0
|
||||
self.start_step = 0
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
@@ -92,102 +45,87 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
self.prompt_cache = PromptEmbedsCache()
|
||||
self.prompt_pairs: list[EncodedPromptPair] = []
|
||||
self.anchor_pairs: list[EncodedAnchor] = []
|
||||
# keep track of prompt chunk size
|
||||
self.prompt_chunk_size = 1
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
|
||||
# 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()
|
||||
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
|
||||
erase_negative = len(target.positive.strip()) == 0
|
||||
enhance_positive = len(target.negative.strip()) == 0
|
||||
# list of neutrals. Can come from file or be empty
|
||||
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
|
||||
|
||||
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:
|
||||
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
|
||||
# remove duplicates
|
||||
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
|
||||
|
||||
if both or erase_negative:
|
||||
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 enhance_positive:
|
||||
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 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
|
||||
),
|
||||
]
|
||||
# encode them
|
||||
cache = encode_prompts_to_cache(
|
||||
prompt_list=prompts_to_cache,
|
||||
sd=self.sd,
|
||||
cache=cache,
|
||||
prompt_tensor_file=self.slider_config.prompt_tensors
|
||||
)
|
||||
|
||||
prompt_pairs = []
|
||||
prompt_batches = []
|
||||
for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
|
||||
for target in self.slider_config.targets:
|
||||
prompt_pair_batch = build_prompt_pair_batch_from_cache(
|
||||
cache=cache,
|
||||
target=target,
|
||||
neutral=neutral,
|
||||
|
||||
)
|
||||
if self.slider_config.batch_full_slide:
|
||||
# concat the prompt pairs
|
||||
# this allows us to run the entire 4 part process in one shot (for slider)
|
||||
self.prompt_chunk_size = 4
|
||||
concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
|
||||
prompt_pairs += [concat_prompt_pair_batch]
|
||||
else:
|
||||
self.prompt_chunk_size = 1
|
||||
# do them one at a time (probably not necessary after new optimizations)
|
||||
prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
|
||||
|
||||
# setup anchors
|
||||
anchor_pairs = []
|
||||
@@ -200,13 +138,26 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
if cache[prompt] == None:
|
||||
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 += [
|
||||
EncodedAnchor(
|
||||
prompt=cache[anchor.prompt],
|
||||
neg_prompt=cache[anchor.neg_prompt],
|
||||
multiplier=anchor.multiplier
|
||||
)
|
||||
concat_anchors(anchor_batch).to('cpu')
|
||||
]
|
||||
if len(anchor_pairs) > 0:
|
||||
self.anchor_pairs = anchor_pairs
|
||||
|
||||
# move to cpu to save vram
|
||||
# We don't need text encoder anymore, but keep it on cpu for sampling
|
||||
@@ -218,48 +169,45 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
self.sd.text_encoder.to("cpu")
|
||||
self.prompt_cache = cache
|
||||
self.prompt_pairs = prompt_pairs
|
||||
self.anchor_pairs = anchor_pairs
|
||||
# self.anchor_pairs = anchor_pairs
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
# get a random pair
|
||||
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
||||
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
|
||||
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
|
||||
# get a random resolution
|
||||
height, width = self.slider_config.resolutions[
|
||||
torch.randint(0, len(self.slider_config.resolutions), (1,)).item()
|
||||
]
|
||||
if self.train_config.gradient_checkpointing:
|
||||
# may get disabled elsewhere
|
||||
self.sd.unet.enable_gradient_checkpointing()
|
||||
|
||||
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.predict_noise(
|
||||
def get_noise_pred(neg, pos, gs, cts, dn):
|
||||
return self.sd.predict_noise(
|
||||
latents=dn,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
p, # unconditional
|
||||
n, # positive
|
||||
neg, # negative prompt
|
||||
pos, # positive prompt
|
||||
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
|
||||
@@ -272,10 +220,15 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
).item()
|
||||
|
||||
# for a complete slider, the batch size is 4 to begin with now
|
||||
true_batch_size = prompt_pair.target_class.text_embeds.shape[0] * self.train_config.batch_size
|
||||
|
||||
# get noise
|
||||
noise = self.get_latent_noise(
|
||||
noise = self.sd.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
batch_size=true_batch_size,
|
||||
noise_offset=self.train_config.noise_offset,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
# get latents
|
||||
@@ -284,12 +237,13 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
self.network.multiplier = multiplier
|
||||
denoised_latents = self.diffuse_some_steps(
|
||||
# pass the multiplier list to the network
|
||||
self.network.multiplier = prompt_pair.multiplier_list
|
||||
denoised_latents = self.sd.diffuse_some_steps(
|
||||
latents, # pass simple noise latents
|
||||
train_tools.concat_prompt_embeddings(
|
||||
positive, # unconditional
|
||||
target_class, # target
|
||||
prompt_pair.positive_target, # unconditional
|
||||
prompt_pair.target_class, # target
|
||||
self.train_config.batch_size,
|
||||
),
|
||||
start_timesteps=0,
|
||||
@@ -297,100 +251,183 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
guidance_scale=3,
|
||||
)
|
||||
|
||||
# split the latents into out prompt pair chunks
|
||||
denoised_latent_chunks = torch.chunk(denoised_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
noise_scheduler.set_timesteps(1000)
|
||||
|
||||
current_timestep = noise_scheduler.timesteps[
|
||||
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
||||
]
|
||||
|
||||
# flush() # 4.2GB to 3GB on 512x512
|
||||
|
||||
# 4.20 GB RAM for 512x512
|
||||
positive_latents = get_noise_pred(
|
||||
positive, negative, 1, current_timestep, denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
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(
|
||||
positive, neutral, 1, current_timestep, denoised_latents
|
||||
).to("cpu", dtype=torch.float32)
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
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(
|
||||
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,
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.positive_target, # positive prompt
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
)
|
||||
erase = prompt_pair.action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE
|
||||
guidance_scale = 1.0
|
||||
unconditional_latents.requires_grad = False
|
||||
unconditional_latents_chunks = torch.chunk(unconditional_latents, self.prompt_chunk_size, dim=0)
|
||||
|
||||
offset = guidance_scale * (positive_latents - unconditional_latents)
|
||||
flush() # 4.2GB to 3GB on 512x512
|
||||
|
||||
offset_neutral = neutral_latents
|
||||
if erase:
|
||||
offset_neutral -= offset
|
||||
else:
|
||||
# enhance
|
||||
offset_neutral += offset
|
||||
# 4.20 GB RAM for 512x512
|
||||
anchor_loss_float = None
|
||||
if len(self.anchor_pairs) > 0:
|
||||
with torch.no_grad():
|
||||
# get a random anchor pair
|
||||
anchor: EncodedAnchor = self.anchor_pairs[
|
||||
torch.randint(0, len(self.anchor_pairs), (1,)).item()
|
||||
]
|
||||
anchor.to(self.device_torch, dtype=dtype)
|
||||
|
||||
loss = loss_function(
|
||||
target_latents,
|
||||
offset_neutral,
|
||||
) * weight
|
||||
# first we get the target prediction without network active
|
||||
anchor_target_noise = get_noise_pred(
|
||||
anchor.neg_prompt, anchor.prompt, 1, current_timestep, denoised_latents
|
||||
# ).to("cpu", dtype=torch.float32)
|
||||
).requires_grad_(False)
|
||||
|
||||
loss_slide = loss.item()
|
||||
# to save vram, we will run these through separately while tracking grads
|
||||
# otherwise it consumes a ton of vram and this isn't our speed bottleneck
|
||||
anchor_chunks = split_anchors(anchor, self.prompt_chunk_size)
|
||||
anchor_target_noise_chunks = torch.chunk(anchor_target_noise, self.prompt_chunk_size, dim=0)
|
||||
assert len(anchor_chunks) == len(denoised_latent_chunks)
|
||||
|
||||
if anchor_loss is not None:
|
||||
loss += anchor_loss
|
||||
# 4.32 GB RAM for 512x512
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
anchor_float_losses = []
|
||||
for anchor_chunk, denoised_latent_chunk, anchor_target_noise_chunk in zip(
|
||||
anchor_chunks, denoised_latent_chunks, anchor_target_noise_chunks
|
||||
):
|
||||
self.network.multiplier = anchor_chunk.multiplier_list
|
||||
|
||||
loss_float = loss.item()
|
||||
anchor_pred_noise = get_noise_pred(
|
||||
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
|
||||
)
|
||||
# 9.42 GB RAM for 512x512 -> 4.20 GB RAM for 512x512 with new grad_checkpointing
|
||||
anchor_loss = loss_function(
|
||||
anchor_target_noise_chunk,
|
||||
anchor_pred_noise,
|
||||
)
|
||||
anchor_float_losses.append(anchor_loss.item())
|
||||
# compute anchor loss gradients
|
||||
# we will accumulate them later
|
||||
# this saves a ton of memory doing them separately
|
||||
anchor_loss.backward()
|
||||
del anchor_pred_noise
|
||||
del anchor_target_noise_chunk
|
||||
del anchor_loss
|
||||
flush()
|
||||
|
||||
loss = loss.to(self.device_torch)
|
||||
anchor_loss_float = sum(anchor_float_losses) / len(anchor_float_losses)
|
||||
del anchor_chunks
|
||||
del anchor_target_noise_chunks
|
||||
del anchor_target_noise
|
||||
# move anchor back to cpu
|
||||
anchor.to("cpu")
|
||||
flush()
|
||||
|
||||
prompt_pair_chunks = split_prompt_pairs(prompt_pair, self.prompt_chunk_size)
|
||||
assert len(prompt_pair_chunks) == len(denoised_latent_chunks)
|
||||
# 3.28 GB RAM for 512x512
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
loss_list = []
|
||||
for prompt_pair_chunk, \
|
||||
denoised_latent_chunk, \
|
||||
positive_latents_chunk, \
|
||||
neutral_latents_chunk, \
|
||||
unconditional_latents_chunk \
|
||||
in zip(
|
||||
prompt_pair_chunks,
|
||||
denoised_latent_chunks,
|
||||
positive_latents_chunks,
|
||||
neutral_latents_chunks,
|
||||
unconditional_latents_chunks,
|
||||
):
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list
|
||||
target_latents = get_noise_pred(
|
||||
prompt_pair_chunk.positive_target,
|
||||
prompt_pair_chunk.target_class,
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latent_chunk
|
||||
)
|
||||
|
||||
guidance_scale = 1.0
|
||||
|
||||
offset = guidance_scale * (positive_latents_chunk - unconditional_latents_chunk)
|
||||
|
||||
# make offset multiplier based on actions
|
||||
offset_multiplier_list = []
|
||||
for action in prompt_pair_chunk.action_list:
|
||||
if action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE:
|
||||
offset_multiplier_list += [-1.0]
|
||||
elif action == ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE:
|
||||
offset_multiplier_list += [1.0]
|
||||
|
||||
offset_multiplier = torch.tensor(offset_multiplier_list).to(offset.device, dtype=offset.dtype)
|
||||
# make offset multiplier match rank of offset
|
||||
offset_multiplier = offset_multiplier.view(offset.shape[0], 1, 1, 1)
|
||||
offset *= offset_multiplier
|
||||
|
||||
offset_neutral = neutral_latents_chunk
|
||||
# offsets are already adjusted on a per-batch basis
|
||||
offset_neutral += offset
|
||||
|
||||
# 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
|
||||
loss = loss_function(
|
||||
target_latents,
|
||||
offset_neutral,
|
||||
) * prompt_pair_chunk.weight
|
||||
|
||||
loss.backward()
|
||||
loss_list.append(loss.item())
|
||||
del target_latents
|
||||
del offset_neutral
|
||||
del loss
|
||||
flush()
|
||||
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
loss_float = sum(loss_list) / len(loss_list)
|
||||
if anchor_loss_float is not None:
|
||||
loss_float += anchor_loss_float
|
||||
|
||||
del (
|
||||
positive_latents,
|
||||
neutral_latents,
|
||||
unconditional_latents,
|
||||
target_latents,
|
||||
latents,
|
||||
latents
|
||||
)
|
||||
# move back to cpu
|
||||
prompt_pair.to("cpu")
|
||||
flush()
|
||||
|
||||
# reset network
|
||||
@@ -399,9 +436,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss_float},
|
||||
)
|
||||
if anchor_loss is not None:
|
||||
loss_dict['sl_l'] = loss_slide
|
||||
loss_dict['an_l'] = anchor_loss.item()
|
||||
if anchor_loss_float is not None:
|
||||
loss_dict['sl_l'] = loss_float
|
||||
loss_dict['an_l'] = anchor_loss_float
|
||||
|
||||
return loss_dict
|
||||
# end hook_train_loop
|
||||
|
||||
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
|
||||
import time
|
||||
import numpy as np
|
||||
from .models.vgg19_critic import Critic
|
||||
|
||||
IMAGE_TRANSFORMS = transforms.Compose(
|
||||
[
|
||||
@@ -37,145 +38,6 @@ def unnormalize(tensor):
|
||||
return (tensor / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
|
||||
class Critic:
|
||||
process: 'TrainVAEProcess'
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
learning_rate=1e-5,
|
||||
device='cpu',
|
||||
optimizer='adam',
|
||||
num_critic_per_gen=1,
|
||||
dtype='float32',
|
||||
lambda_gp=10,
|
||||
start_step=0,
|
||||
warmup_steps=1000,
|
||||
process=None,
|
||||
optimizer_params=None,
|
||||
):
|
||||
self.learning_rate = learning_rate
|
||||
self.device = device
|
||||
self.optimizer_type = optimizer
|
||||
self.num_critic_per_gen = num_critic_per_gen
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.process = process
|
||||
self.model = None
|
||||
self.optimizer = None
|
||||
self.scheduler = None
|
||||
self.warmup_steps = warmup_steps
|
||||
self.start_step = start_step
|
||||
self.lambda_gp = lambda_gp
|
||||
|
||||
if optimizer_params is None:
|
||||
optimizer_params = {}
|
||||
self.optimizer_params = optimizer_params
|
||||
self.print = self.process.print
|
||||
print(f" Critic config: {self.__dict__}")
|
||||
|
||||
def setup(self):
|
||||
from .models.vgg19_critic import Vgg19Critic
|
||||
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
|
||||
self.load_weights()
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
params = self.model.parameters()
|
||||
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
|
||||
self.optimizer,
|
||||
total_iters=self.process.max_steps * self.num_critic_per_gen,
|
||||
factor=1,
|
||||
verbose=False
|
||||
)
|
||||
|
||||
def load_weights(self):
|
||||
path_to_load = None
|
||||
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
|
||||
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
path_to_load = latest_file
|
||||
else:
|
||||
self.print(f" - No checkpoint found, starting from scratch")
|
||||
if path_to_load:
|
||||
self.model.load_state_dict(load_file(path_to_load))
|
||||
|
||||
def save(self, step=None):
|
||||
self.process.update_training_metadata()
|
||||
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zeropad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
|
||||
save_file(self.model.state_dict(), save_path, save_meta)
|
||||
self.print(f"Saved critic to {save_path}")
|
||||
|
||||
def get_critic_loss(self, vgg_output):
|
||||
if self.start_step > self.process.step_num:
|
||||
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
|
||||
|
||||
warmup_scaler = 1.0
|
||||
# we need a warmup when we come on of 1000 steps
|
||||
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
|
||||
if self.process.step_num < self.start_step + self.warmup_steps:
|
||||
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
|
||||
# set model to not train for generator loss
|
||||
self.model.eval()
|
||||
self.model.requires_grad_(False)
|
||||
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
|
||||
|
||||
# run model
|
||||
stacked_output = self.model(vgg_pred)
|
||||
|
||||
return (-torch.mean(stacked_output)) * warmup_scaler
|
||||
|
||||
def step(self, vgg_output):
|
||||
|
||||
# train critic here
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
|
||||
critic_losses = []
|
||||
for i in range(self.num_critic_per_gen):
|
||||
inputs = vgg_output.detach()
|
||||
inputs = inputs.to(self.device, dtype=self.torch_dtype)
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
|
||||
|
||||
stacked_output = self.model(inputs)
|
||||
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
|
||||
|
||||
# Compute gradient penalty
|
||||
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
|
||||
|
||||
# Compute WGAN-GP critic loss
|
||||
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
|
||||
critic_loss.backward()
|
||||
self.optimizer.zero_grad()
|
||||
self.optimizer.step()
|
||||
self.scheduler.step()
|
||||
critic_losses.append(critic_loss.item())
|
||||
|
||||
# avg loss
|
||||
loss = np.mean(critic_losses)
|
||||
return loss
|
||||
|
||||
def get_lr(self):
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
learning_rate = (
|
||||
self.optimizer.param_groups[0]["d"] *
|
||||
self.optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
learning_rate = self.optimizer.param_groups[0]['lr']
|
||||
|
||||
return learning_rate
|
||||
|
||||
|
||||
class TrainVAEProcess(BaseTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
@@ -6,5 +6,11 @@ from .BaseTrainProcess import BaseTrainProcess
|
||||
from .TrainVAEProcess import TrainVAEProcess
|
||||
from .BaseMergeProcess import BaseMergeProcess
|
||||
from .TrainSliderProcess import TrainSliderProcess
|
||||
from .TrainSliderProcessOld import TrainSliderProcessOld
|
||||
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.nn as nn
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
from toolkit.losses import get_gradient_penalty
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.optimizer import get_optimizer
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
|
||||
class MeanReduce(nn.Module):
|
||||
@@ -36,3 +48,147 @@ class Vgg19Critic(nn.Module):
|
||||
|
||||
def forward(self, inputs):
|
||||
return self.main(inputs)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from jobs.process.TrainVAEProcess import TrainVAEProcess
|
||||
from jobs.process.TrainESRGANProcess import TrainESRGANProcess
|
||||
|
||||
|
||||
class Critic:
|
||||
process: Union['TrainVAEProcess', 'TrainESRGANProcess']
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
learning_rate=1e-5,
|
||||
device='cpu',
|
||||
optimizer='adam',
|
||||
num_critic_per_gen=1,
|
||||
dtype='float32',
|
||||
lambda_gp=10,
|
||||
start_step=0,
|
||||
warmup_steps=1000,
|
||||
process=None,
|
||||
optimizer_params=None,
|
||||
):
|
||||
self.learning_rate = learning_rate
|
||||
self.device = device
|
||||
self.optimizer_type = optimizer
|
||||
self.num_critic_per_gen = num_critic_per_gen
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.process = process
|
||||
self.model = None
|
||||
self.optimizer = None
|
||||
self.scheduler = None
|
||||
self.warmup_steps = warmup_steps
|
||||
self.start_step = start_step
|
||||
self.lambda_gp = lambda_gp
|
||||
|
||||
if optimizer_params is None:
|
||||
optimizer_params = {}
|
||||
self.optimizer_params = optimizer_params
|
||||
self.print = self.process.print
|
||||
print(f" Critic config: {self.__dict__}")
|
||||
|
||||
def setup(self):
|
||||
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
|
||||
self.load_weights()
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
params = self.model.parameters()
|
||||
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
|
||||
self.optimizer,
|
||||
total_iters=self.process.max_steps * self.num_critic_per_gen,
|
||||
factor=1,
|
||||
verbose=False
|
||||
)
|
||||
|
||||
def load_weights(self):
|
||||
path_to_load = None
|
||||
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
|
||||
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
path_to_load = latest_file
|
||||
else:
|
||||
self.print(f" - No checkpoint found, starting from scratch")
|
||||
if path_to_load:
|
||||
self.model.load_state_dict(load_file(path_to_load))
|
||||
|
||||
def save(self, step=None):
|
||||
self.process.update_training_metadata()
|
||||
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zeropad 9 digits
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
|
||||
save_file(self.model.state_dict(), save_path, save_meta)
|
||||
self.print(f"Saved critic to {save_path}")
|
||||
|
||||
def get_critic_loss(self, vgg_output):
|
||||
if self.start_step > self.process.step_num:
|
||||
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
|
||||
|
||||
warmup_scaler = 1.0
|
||||
# we need a warmup when we come on of 1000 steps
|
||||
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
|
||||
if self.process.step_num < self.start_step + self.warmup_steps:
|
||||
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
|
||||
# set model to not train for generator loss
|
||||
self.model.eval()
|
||||
self.model.requires_grad_(False)
|
||||
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
|
||||
|
||||
# run model
|
||||
stacked_output = self.model(vgg_pred)
|
||||
|
||||
return (-torch.mean(stacked_output)) * warmup_scaler
|
||||
|
||||
def step(self, vgg_output):
|
||||
|
||||
# train critic here
|
||||
self.model.train()
|
||||
self.model.requires_grad_(True)
|
||||
|
||||
critic_losses = []
|
||||
for i in range(self.num_critic_per_gen):
|
||||
inputs = vgg_output.detach()
|
||||
inputs = inputs.to(self.device, dtype=self.torch_dtype)
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
|
||||
|
||||
stacked_output = self.model(inputs)
|
||||
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
|
||||
|
||||
# Compute gradient penalty
|
||||
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
|
||||
|
||||
# Compute WGAN-GP critic loss
|
||||
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
|
||||
critic_loss.backward()
|
||||
self.optimizer.zero_grad()
|
||||
self.optimizer.step()
|
||||
self.scheduler.step()
|
||||
critic_losses.append(critic_loss.item())
|
||||
|
||||
# avg loss
|
||||
loss = np.mean(critic_losses)
|
||||
return loss
|
||||
|
||||
def get_lr(self):
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
learning_rate = (
|
||||
self.optimizer.param_groups[0]["d"] *
|
||||
self.optimizer.param_groups[0]["lr"]
|
||||
)
|
||||
else:
|
||||
learning_rate = self.optimizer.param_groups[0]['lr']
|
||||
|
||||
return learning_rate
|
||||
|
||||
|
||||
338
notebooks/SliderTraining.ipynb
Normal file
338
notebooks/SliderTraining.ipynb
Normal file
@@ -0,0 +1,338 @@
|
||||
{
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0,
|
||||
"metadata": {
|
||||
"colab": {
|
||||
"provenance": [],
|
||||
"machine_shape": "hm",
|
||||
"gpuType": "V100"
|
||||
},
|
||||
"kernelspec": {
|
||||
"name": "python3",
|
||||
"display_name": "Python 3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
},
|
||||
"accelerator": "GPU"
|
||||
},
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## Slider Training\n",
|
||||
"\n",
|
||||
"This is a quick colab demo for training sliders like can be found in my CivitAI profile https://civitai.com/user/Ostris/models . I will work on making it more user friendly, but for now, it will get you started."
|
||||
],
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image"
|
||||
],
|
||||
"metadata": {
|
||||
"collapsed": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
],
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict({\n",
|
||||
" # This is the config I use on my sliders, It is solid and tested\n",
|
||||
" 'job': 'train',\n",
|
||||
" 'config': {\n",
|
||||
" # the name will be used to create a folder in the output folder\n",
|
||||
" # it will also replace any [name] token in the rest of this config\n",
|
||||
" 'name': 'detail_slider_v1',\n",
|
||||
" # folder will be created with name above in folder below\n",
|
||||
" # it can be relative to the project root or absolute\n",
|
||||
" 'training_folder': \"output/LoRA\",\n",
|
||||
" 'device': 'cuda', # cpu, cuda:0, etc\n",
|
||||
" # for tensorboard logging, we will make a subfolder for this job\n",
|
||||
" 'log_dir': \"output/.tensorboard\",\n",
|
||||
" # you can stack processes for other jobs, It is not tested with sliders though\n",
|
||||
" # just use one for now\n",
|
||||
" 'process': [\n",
|
||||
" {\n",
|
||||
" 'type': 'slider', # tells runner to run the slider process\n",
|
||||
" # network is the LoRA network for a slider, I recommend to leave this be\n",
|
||||
" 'network': {\n",
|
||||
" 'type': \"lora\",\n",
|
||||
" # rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good\n",
|
||||
" 'linear': 8, # \"rank\" or \"dim\"\n",
|
||||
" 'linear_alpha': 4, # Do about half of rank \"alpha\"\n",
|
||||
" # 'conv': 4, # for convolutional layers \"locon\"\n",
|
||||
" # 'conv_alpha': 4, # Do about half of conv \"alpha\"\n",
|
||||
" },\n",
|
||||
" # training config\n",
|
||||
" 'train': {\n",
|
||||
" # this is also used in sampling. Stick with ddpm unless you know what you are doing\n",
|
||||
" 'noise_scheduler': \"ddpm\", # or \"ddpm\", \"lms\", \"euler_a\"\n",
|
||||
" # how many steps to train. More is not always better. I rarely go over 1000\n",
|
||||
" 'steps': 100,\n",
|
||||
" # I have had good results with 4e-4 to 1e-4 at 500 steps\n",
|
||||
" 'lr': 2e-4,\n",
|
||||
" # enables gradient checkpoint, saves vram, leave it on\n",
|
||||
" 'gradient_checkpointing': True,\n",
|
||||
" # train the unet. I recommend leaving this true\n",
|
||||
" 'train_unet': True,\n",
|
||||
" # train the text encoder. I don't recommend this unless you have a special use case\n",
|
||||
" # for sliders we are adjusting representation of the concept (unet),\n",
|
||||
" # not the description of it (text encoder)\n",
|
||||
" 'train_text_encoder': False,\n",
|
||||
"\n",
|
||||
" # just leave unless you know what you are doing\n",
|
||||
" # also supports \"dadaptation\" but set lr to 1 if you use that,\n",
|
||||
" # but it learns too fast and I don't recommend it\n",
|
||||
" 'optimizer': \"adamw\",\n",
|
||||
" # only constant for now\n",
|
||||
" 'lr_scheduler': \"constant\",\n",
|
||||
" # we randomly denoise random num of steps form 1 to this number\n",
|
||||
" # while training. Just leave it\n",
|
||||
" 'max_denoising_steps': 40,\n",
|
||||
" # works great at 1. I do 1 even with my 4090.\n",
|
||||
" # higher may not work right with newer single batch stacking code anyway\n",
|
||||
" 'batch_size': 1,\n",
|
||||
" # bf16 works best if your GPU supports it (modern)\n",
|
||||
" 'dtype': 'bf16', # fp32, bf16, fp16\n",
|
||||
" # I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX\n",
|
||||
" # although, the way we train sliders is comparative, so it probably won't work anyway\n",
|
||||
" 'noise_offset': 0.0,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # the model to train the LoRA network on\n",
|
||||
" 'model': {\n",
|
||||
" # name_or_path can be a hugging face name, local path or url to model\n",
|
||||
" # on civit ai with or without modelVersionId. They will be cached in /model folder\n",
|
||||
" # epicRealisim v5\n",
|
||||
" 'name_or_path': \"https://civitai.com/models/25694?modelVersionId=134065\",\n",
|
||||
" 'is_v2': False, # for v2 models\n",
|
||||
" 'is_v_pred': False, # for v-prediction models (most v2 models)\n",
|
||||
" # has some issues with the dual text encoder and the way we train sliders\n",
|
||||
" # it works bit weights need to probably be higher to see it.\n",
|
||||
" 'is_xl': False, # for SDXL models\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # saving config\n",
|
||||
" 'save': {\n",
|
||||
" 'dtype': 'float16', # precision to save. I recommend float16\n",
|
||||
" 'save_every': 50, # save every this many steps\n",
|
||||
" # this will remove step counts more than this number\n",
|
||||
" # allows you to save more often in case of a crash without filling up your drive\n",
|
||||
" 'max_step_saves_to_keep': 2,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # sampling config\n",
|
||||
" 'sample': {\n",
|
||||
" # must match train.noise_scheduler, this is not used here\n",
|
||||
" # but may be in future and in other processes\n",
|
||||
" 'sampler': \"ddpm\",\n",
|
||||
" # sample every this many steps\n",
|
||||
" 'sample_every': 20,\n",
|
||||
" # image size\n",
|
||||
" 'width': 512,\n",
|
||||
" 'height': 512,\n",
|
||||
" # prompts to use for sampling. Do as many as you want, but it slows down training\n",
|
||||
" # pick ones that will best represent the concept you are trying to adjust\n",
|
||||
" # allows some flags after the prompt\n",
|
||||
" # --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive\n",
|
||||
" # slide are good tests. will inherit sample.network_multiplier if not set\n",
|
||||
" # --n [string] # negative prompt, will inherit sample.neg if not set\n",
|
||||
" # Only 75 tokens allowed currently\n",
|
||||
" # I like to do a wide positive and negative spread so I can see a good range and stop\n",
|
||||
" # early if the network is braking down\n",
|
||||
" 'prompts': [\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3\",\n",
|
||||
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5\",\n",
|
||||
" \"a golden retriever sitting on a leather couch, --m -5\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m -3\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m 3\",\n",
|
||||
" \"a golden retriever sitting on a leather couch --m 5\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3\",\n",
|
||||
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5\",\n",
|
||||
" ],\n",
|
||||
" # negative prompt used on all prompts above as default if they don't have one\n",
|
||||
" 'neg': \"cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome\",\n",
|
||||
" # seed for sampling. 42 is the answer for everything\n",
|
||||
" 'seed': 42,\n",
|
||||
" # walks the seed so s1 is 42, s2 is 43, s3 is 44, etc\n",
|
||||
" # will start over on next sample_every so s1 is always seed\n",
|
||||
" # works well if you use same prompt but want different results\n",
|
||||
" 'walk_seed': False,\n",
|
||||
" # cfg scale (4 to 10 is good)\n",
|
||||
" 'guidance_scale': 7,\n",
|
||||
" # sampler steps (20 to 30 is good)\n",
|
||||
" 'sample_steps': 20,\n",
|
||||
" # default network multiplier for all prompts\n",
|
||||
" # since we are training a slider, I recommend overriding this with --m [number]\n",
|
||||
" # in the prompts above to get both sides of the slider\n",
|
||||
" 'network_multiplier': 1.0,\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # logging information\n",
|
||||
" 'logging': {\n",
|
||||
" 'log_every': 10, # log every this many steps\n",
|
||||
" 'use_wandb': False, # not supported yet\n",
|
||||
" 'verbose': False, # probably done need unless you are debugging\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # slider training config, best for last\n",
|
||||
" 'slider': {\n",
|
||||
" # resolutions to train on. [ width, height ]. This is less important for sliders\n",
|
||||
" # as we are not teaching the model anything it doesn't already know\n",
|
||||
" # but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1\n",
|
||||
" # and [ 1024, 1024 ] for sd_xl\n",
|
||||
" # you can do as many as you want here\n",
|
||||
" 'resolutions': [\n",
|
||||
" [512, 512],\n",
|
||||
" # [ 512, 768 ]\n",
|
||||
" # [ 768, 768 ]\n",
|
||||
" ],\n",
|
||||
" # slider training uses 4 combined steps for a single round. This will do it in one gradient\n",
|
||||
" # step. It is highly optimized and shouldn't take anymore vram than doing without it,\n",
|
||||
" # since we break down batches for gradient accumulation now. so just leave it on.\n",
|
||||
" 'batch_full_slide': True,\n",
|
||||
" # These are the concepts to train on. You can do as many as you want here,\n",
|
||||
" # but they can conflict outweigh each other. Other than experimenting, I recommend\n",
|
||||
" # just doing one for good results\n",
|
||||
" 'targets': [\n",
|
||||
" # target_class is the base concept we are adjusting the representation of\n",
|
||||
" # for example, if we are adjusting the representation of a person, we would use \"person\"\n",
|
||||
" # if we are adjusting the representation of a cat, we would use \"cat\" It is not\n",
|
||||
" # a keyword necessarily but what the model understands the concept to represent.\n",
|
||||
" # \"person\" will affect men, women, children, etc but will not affect cats, dogs, etc\n",
|
||||
" # it is the models base general understanding of the concept and everything it represents\n",
|
||||
" # you can leave it blank to affect everything. In this example, we are adjusting\n",
|
||||
" # detail, so we will leave it blank to affect everything\n",
|
||||
" {\n",
|
||||
" 'target_class': \"\",\n",
|
||||
" # positive is the prompt for the positive side of the slider.\n",
|
||||
" # It is the concept that will be excited and amplified in the model when we slide the slider\n",
|
||||
" # to the positive side and forgotten / inverted when we slide\n",
|
||||
" # the slider to the negative side. It is generally best to include the target_class in\n",
|
||||
" # the prompt. You want it to be the extreme of what you want to train on. For example,\n",
|
||||
" # if you want to train on fat people, you would use \"an extremely fat, morbidly obese person\"\n",
|
||||
" # as the prompt. Not just \"fat person\"\n",
|
||||
" # max 75 tokens for now\n",
|
||||
" 'positive': \"high detail, 8k, intricate, detailed, high resolution, high res, high quality\",\n",
|
||||
" # negative is the prompt for the negative side of the slider and works the same as positive\n",
|
||||
" # it does not necessarily work the same as a negative prompt when generating images\n",
|
||||
" # these need to be polar opposites.\n",
|
||||
" # max 76 tokens for now\n",
|
||||
" 'negative': \"blurry, boring, fuzzy, low detail, low resolution, low res, low quality\",\n",
|
||||
" # the loss for this target is multiplied by this number.\n",
|
||||
" # if you are doing more than one target it may be good to set less important ones\n",
|
||||
" # to a lower number like 0.1 so they don't outweigh the primary target\n",
|
||||
" 'weight': 1.0,\n",
|
||||
" },\n",
|
||||
" ],\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
" ]\n",
|
||||
" },\n",
|
||||
"\n",
|
||||
" # You can put any information you want here, and it will be saved in the model.\n",
|
||||
" # The below is an example, but you can put your grocery list in it if you want.\n",
|
||||
" # It is saved in the model so be aware of that. The software will include this\n",
|
||||
" # plus some other information for you automatically\n",
|
||||
" 'meta': {\n",
|
||||
" # [name] gets replaced with the name above\n",
|
||||
" 'name': \"[name]\",\n",
|
||||
" 'version': '1.0',\n",
|
||||
" # 'creator': {\n",
|
||||
" # 'name': 'your name',\n",
|
||||
" # 'email': 'your@gmail.com',\n",
|
||||
" # 'website': 'https://your.website'\n",
|
||||
" # }\n",
|
||||
" }\n",
|
||||
"})\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. Ill update soon."
|
||||
],
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
}
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"execution_count": null,
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
],
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -5,7 +5,6 @@ diffusers
|
||||
transformers
|
||||
lycoris_lora
|
||||
flatten_json
|
||||
accelerator
|
||||
pyyaml
|
||||
oyaml
|
||||
tensorboard
|
||||
@@ -14,4 +13,6 @@ invisible-watermark
|
||||
einops
|
||||
accelerate
|
||||
toml
|
||||
albumentations
|
||||
albumentations
|
||||
pydantic
|
||||
omegaconf
|
||||
|
||||
12
run.py
12
run.py
@@ -1,5 +1,7 @@
|
||||
import os
|
||||
import sys
|
||||
from typing import Union, OrderedDict
|
||||
|
||||
sys.path.insert(0, os.getcwd())
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
@@ -36,6 +38,14 @@ def main():
|
||||
action='store_true',
|
||||
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()
|
||||
|
||||
config_file_list = args.config_file_list
|
||||
@@ -49,7 +59,7 @@ def main():
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
job = get_job(config_file)
|
||||
job = get_job(config_file, args.name)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
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 json
|
||||
from typing import Union
|
||||
|
||||
import oyaml as yaml
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
@@ -15,22 +17,24 @@ def get_cwd_abs_path(path):
|
||||
return path
|
||||
|
||||
|
||||
def preprocess_config(config: OrderedDict):
|
||||
def preprocess_config(config: OrderedDict, name: str = None):
|
||||
if "job" not in config:
|
||||
raise ValueError("config file must have a job key")
|
||||
if "config" not in config:
|
||||
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")
|
||||
# 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 = config_string.replace("[name]", name)
|
||||
config = json.loads(config_string, object_pairs_hook=OrderedDict)
|
||||
return config
|
||||
|
||||
|
||||
|
||||
# Fixes issue where yaml doesnt load exponents correctly
|
||||
fixed_loader = yaml.SafeLoader
|
||||
fixed_loader.add_implicit_resolver(
|
||||
@@ -44,7 +48,18 @@ fixed_loader.add_implicit_resolver(
|
||||
|\\.(?:nan|NaN|NAN))$''', re.X),
|
||||
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
|
||||
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
|
||||
@@ -67,12 +82,12 @@ def get_config(config_file_path):
|
||||
|
||||
# if we found it, check if it is a json or yaml file
|
||||
if real_config_path.endswith('.json') or real_config_path.endswith('.jsonc'):
|
||||
with open(real_config_path, 'r') as f:
|
||||
with open(real_config_path, 'r', encoding='utf-8') as f:
|
||||
config = json.load(f, object_pairs_hook=OrderedDict)
|
||||
elif real_config_path.endswith('.yaml') or real_config_path.endswith('.yml'):
|
||||
with open(real_config_path, 'r') as f:
|
||||
with open(real_config_path, 'r', encoding='utf-8') as f:
|
||||
config = yaml.load(f, Loader=fixed_loader)
|
||||
else:
|
||||
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
||||
|
||||
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:
|
||||
@@ -27,6 +30,7 @@ class SampleConfig:
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.network_multiplier = kwargs.get('network_multiplier', 1)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
|
||||
|
||||
class NetworkConfig:
|
||||
@@ -35,13 +39,15 @@ class NetworkConfig:
|
||||
rank = kwargs.get('rank', None)
|
||||
linear = kwargs.get('linear', None)
|
||||
if rank is not None:
|
||||
self.rank: int = rank # rank for backward compatibility
|
||||
self.rank: int = rank # rank for backward compatibility
|
||||
self.linear: int = rank
|
||||
elif linear is not None:
|
||||
self.rank: int = linear
|
||||
self.linear: int = linear
|
||||
self.conv: int = kwargs.get('conv', None)
|
||||
self.alpha: float = kwargs.get('alpha', 1.0)
|
||||
self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha)
|
||||
self.conv_alpha: float = kwargs.get('conv_alpha', self.conv)
|
||||
|
||||
|
||||
class TrainConfig:
|
||||
@@ -60,7 +66,7 @@ class TrainConfig:
|
||||
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
||||
self.optimizer_params = kwargs.get('optimizer_params', {})
|
||||
self.skip_first_sample = kwargs.get('skip_first_sample', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True)
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
@@ -69,6 +75,8 @@ class ModelConfig:
|
||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.vae_path: str = kwargs.get('vae_path', None)
|
||||
|
||||
if self.name_or_path is None:
|
||||
raise ValueError('name_or_path must be specified')
|
||||
@@ -99,3 +107,200 @@ class SliderConfig:
|
||||
anchors = [SliderConfigAnchors(**anchor) for anchor in anchors]
|
||||
self.anchors: List[SliderConfigAnchors] = anchors
|
||||
self.resolutions: List[List[int]] = kwargs.get('resolutions', [[512, 512]])
|
||||
self.prompt_file: str = kwargs.get('prompt_file', None)
|
||||
self.prompt_tensors: str = kwargs.get('prompt_tensors', None)
|
||||
self.batch_full_slide: bool = kwargs.get('batch_full_slide', True)
|
||||
|
||||
|
||||
class GenerateImageConfig:
|
||||
def __init__(
|
||||
self,
|
||||
prompt: str = '',
|
||||
prompt_2: Optional[str] = None,
|
||||
width: int = 512,
|
||||
height: int = 512,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 7.5,
|
||||
negative_prompt: str = '',
|
||||
negative_prompt_2: Optional[str] = None,
|
||||
seed: int = -1,
|
||||
network_multiplier: float = 1.0,
|
||||
guidance_rescale: float = 0.0,
|
||||
# the tag [time] will be replaced with milliseconds since epoch
|
||||
output_path: str = None, # full image path
|
||||
output_folder: str = None, # folder to save image in if output_path is not specified
|
||||
output_ext: str = 'png', # extension to save image as if output_path is not specified
|
||||
output_tail: str = '', # tail to add to output filename
|
||||
add_prompt_file: bool = False, # add a prompt file with generated image
|
||||
):
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
self.num_inference_steps: int = num_inference_steps
|
||||
self.guidance_scale: float = guidance_scale
|
||||
self.guidance_rescale: float = guidance_rescale
|
||||
self.prompt: str = prompt
|
||||
self.prompt_2: str = prompt_2
|
||||
self.negative_prompt: str = negative_prompt
|
||||
self.negative_prompt_2: str = negative_prompt_2
|
||||
|
||||
self.output_path: str = output_path
|
||||
self.seed: int = seed
|
||||
if self.seed == -1:
|
||||
# generate random one
|
||||
self.seed = random.randint(0, 2 ** 32 - 1)
|
||||
self.network_multiplier: float = network_multiplier
|
||||
self.output_folder: str = output_folder
|
||||
self.output_ext: str = output_ext
|
||||
self.add_prompt_file: bool = add_prompt_file
|
||||
self.output_tail: str = output_tail
|
||||
self.gen_time: int = int(time.time() * 1000)
|
||||
|
||||
# prompt string will override any settings above
|
||||
self._process_prompt_string()
|
||||
|
||||
# handle dual text encoder prompts if nothing passed
|
||||
if negative_prompt_2 is None:
|
||||
self.negative_prompt_2 = negative_prompt
|
||||
|
||||
if prompt_2 is None:
|
||||
self.prompt_2 = prompt
|
||||
|
||||
# parse prompt paths
|
||||
if self.output_path is None and self.output_folder is None:
|
||||
raise ValueError('output_path or output_folder must be specified')
|
||||
elif self.output_path is not None:
|
||||
self.output_folder = os.path.dirname(self.output_path)
|
||||
self.output_ext = os.path.splitext(self.output_path)[1][1:]
|
||||
self.output_filename_no_ext = os.path.splitext(os.path.basename(self.output_path))[0]
|
||||
|
||||
else:
|
||||
self.output_filename_no_ext = '[time]_[count]'
|
||||
if len(self.output_tail) > 0:
|
||||
self.output_filename_no_ext += '_' + self.output_tail
|
||||
self.output_path = os.path.join(self.output_folder, self.output_filename_no_ext + '.' + self.output_ext)
|
||||
|
||||
# adjust height
|
||||
self.height = max(64, self.height - self.height % 8) # round to divisible by 8
|
||||
self.width = max(64, self.width - self.width % 8) # round to divisible by 8
|
||||
|
||||
def set_gen_time(self, gen_time: int = None):
|
||||
if gen_time is not None:
|
||||
self.gen_time = gen_time
|
||||
else:
|
||||
self.gen_time = int(time.time() * 1000)
|
||||
|
||||
def _get_path_no_ext(self, count: int = 0, max_count=0):
|
||||
# zero pad count
|
||||
count_str = str(count).zfill(len(str(max_count)))
|
||||
# replace [time] with gen time
|
||||
filename = self.output_filename_no_ext.replace('[time]', str(self.gen_time))
|
||||
# replace [count] with count
|
||||
filename = filename.replace('[count]', count_str)
|
||||
return filename
|
||||
|
||||
def get_image_path(self, count: int = 0, max_count=0):
|
||||
filename = self._get_path_no_ext(count, max_count)
|
||||
filename += '.' + self.output_ext
|
||||
# join with folder
|
||||
return os.path.join(self.output_folder, filename)
|
||||
|
||||
def get_prompt_path(self, count: int = 0, max_count=0):
|
||||
filename = self._get_path_no_ext(count, max_count)
|
||||
filename += '.txt'
|
||||
# join with folder
|
||||
return os.path.join(self.output_folder, filename)
|
||||
|
||||
def save_image(self, image, count: int = 0, max_count=0):
|
||||
# make parent dirs
|
||||
os.makedirs(self.output_folder, exist_ok=True)
|
||||
self.set_gen_time()
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
|
||||
def save_prompt_file(self, count: int = 0, max_count=0):
|
||||
# save prompt file
|
||||
with open(self.get_prompt_path(count, max_count), 'w') as f:
|
||||
prompt = self.prompt
|
||||
if self.prompt_2 is not None:
|
||||
prompt += ' --p2 ' + self.prompt_2
|
||||
if self.negative_prompt is not None:
|
||||
prompt += ' --n ' + self.negative_prompt
|
||||
if self.negative_prompt_2 is not None:
|
||||
prompt += ' --n2 ' + self.negative_prompt_2
|
||||
prompt += ' --w ' + str(self.width)
|
||||
prompt += ' --h ' + str(self.height)
|
||||
prompt += ' --seed ' + str(self.seed)
|
||||
prompt += ' --cfg ' + str(self.guidance_scale)
|
||||
prompt += ' --steps ' + str(self.num_inference_steps)
|
||||
prompt += ' --m ' + str(self.network_multiplier)
|
||||
prompt += ' --gr ' + str(self.guidance_rescale)
|
||||
|
||||
# get gen info
|
||||
f.write(self.prompt)
|
||||
|
||||
def _process_prompt_string(self):
|
||||
# we will try to support all sd-scripts where we can
|
||||
|
||||
# FROM SD-SCRIPTS
|
||||
# --n Treat everything until the next option as a negative prompt.
|
||||
# --w Specify the width of the generated image.
|
||||
# --h Specify the height of the generated image.
|
||||
# --d Specify the seed for the generated image.
|
||||
# --l Specify the CFG scale for the generated image.
|
||||
# --s Specify the number of steps during generation.
|
||||
|
||||
# OURS and some QOL additions
|
||||
# --m Specify the network multiplier for the generated image.
|
||||
# --p2 Prompt for the second text encoder (SDXL only)
|
||||
# --n2 Negative prompt for the second text encoder (SDXL only)
|
||||
# --gr Specify the guidance rescale for the generated image (SDXL only)
|
||||
|
||||
# --seed Specify the seed for the generated image same as --d
|
||||
# --cfg Specify the CFG scale for the generated image same as --l
|
||||
# --steps Specify the number of steps during generation same as --s
|
||||
# --network_multiplier Specify the network multiplier for the generated image same as --m
|
||||
|
||||
# process prompt string and update values if it has some
|
||||
if self.prompt is not None and len(self.prompt) > 0:
|
||||
# process prompt string
|
||||
prompt = self.prompt
|
||||
prompt = prompt.strip()
|
||||
p_split = prompt.split('--')
|
||||
self.prompt = p_split[0].strip()
|
||||
|
||||
if len(p_split) > 1:
|
||||
for split in p_split[1:]:
|
||||
# allows multi char flags
|
||||
flag = split.split(' ')[0].strip()
|
||||
content = split[len(flag):].strip()
|
||||
if flag == 'p2':
|
||||
self.prompt_2 = content
|
||||
elif flag == 'n':
|
||||
self.negative_prompt = content
|
||||
elif flag == 'n2':
|
||||
self.negative_prompt_2 = content
|
||||
elif flag == 'w':
|
||||
self.width = int(content)
|
||||
elif flag == 'h':
|
||||
self.height = int(content)
|
||||
elif flag == 'd':
|
||||
self.seed = int(content)
|
||||
elif flag == 'seed':
|
||||
self.seed = int(content)
|
||||
elif flag == 'l':
|
||||
self.guidance_scale = float(content)
|
||||
elif flag == 'cfg':
|
||||
self.guidance_scale = float(content)
|
||||
elif flag == 's':
|
||||
self.num_inference_steps = int(content)
|
||||
elif flag == 'steps':
|
||||
self.num_inference_steps = int(content)
|
||||
elif flag == 'm':
|
||||
self.network_multiplier = float(content)
|
||||
elif flag == 'network_multiplier':
|
||||
self.network_multiplier = float(content)
|
||||
elif flag == 'gr':
|
||||
self.guidance_rescale = float(content)
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from torchvision import transforms
|
||||
from torch.utils.data import Dataset
|
||||
from tqdm import tqdm
|
||||
import albumentations as A
|
||||
|
||||
|
||||
class ImageDataset(Dataset):
|
||||
@@ -38,7 +42,7 @@ class ImageDataset(Dataset):
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
@@ -65,7 +69,7 @@ class ImageDataset(Dataset):
|
||||
if self.random_scale and min_img_size > self.resolution:
|
||||
if min_img_size < self.resolution:
|
||||
print(
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={file}")
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={img_path}")
|
||||
scale_size = self.resolution
|
||||
else:
|
||||
scale_size = random.randint(self.resolution, int(min_img_size))
|
||||
@@ -78,3 +82,124 @@ class ImageDataset(Dataset):
|
||||
img = self.transform(img)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
class Augments:
|
||||
def __init__(self, **kwargs):
|
||||
self.method_name = kwargs.get('method', None)
|
||||
self.params = kwargs.get('params', {})
|
||||
|
||||
# convert kwargs enums for cv2
|
||||
for key, value in self.params.items():
|
||||
if isinstance(value, str):
|
||||
# split the string
|
||||
split_string = value.split('.')
|
||||
if len(split_string) == 2 and split_string[0] == 'cv2':
|
||||
if hasattr(cv2, split_string[1]):
|
||||
self.params[key] = getattr(cv2, split_string[1].upper())
|
||||
else:
|
||||
raise ValueError(f"invalid cv2 enum: {split_string[1]}")
|
||||
|
||||
|
||||
class AugmentedImageDataset(ImageDataset):
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
self.augmentations = self.get_config('augmentations', [])
|
||||
self.augmentations = [Augments(**aug) for aug in self.augmentations]
|
||||
|
||||
augmentation_list = []
|
||||
for aug in self.augmentations:
|
||||
# make sure method name is valid
|
||||
assert hasattr(A, aug.method_name), f"invalid augmentation method: {aug.method_name}"
|
||||
# get the method
|
||||
method = getattr(A, aug.method_name)
|
||||
# add the method to the list
|
||||
augmentation_list.append(method(**aug.params))
|
||||
|
||||
self.aug_transform = A.Compose(augmentation_list)
|
||||
self.original_transform = self.transform
|
||||
# replace transform so we get raw pil image
|
||||
self.transform = transforms.Compose([])
|
||||
|
||||
def __getitem__(self, index):
|
||||
# get the original image
|
||||
# image is a PIL image, convert to bgr
|
||||
pil_image = super().__getitem__(index)
|
||||
open_cv_image = np.array(pil_image)
|
||||
# Convert RGB to BGR
|
||||
open_cv_image = open_cv_image[:, :, ::-1].copy()
|
||||
|
||||
# apply augmentations
|
||||
augmented = self.aug_transform(image=open_cv_image)["image"]
|
||||
|
||||
# convert back to RGB tensor
|
||||
augmented = cv2.cvtColor(augmented, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# convert to PIL image
|
||||
augmented = Image.fromarray(augmented)
|
||||
|
||||
# return both # return image as 0 - 1 tensor
|
||||
return transforms.ToTensor()(pil_image), transforms.ToTensor()(augmented)
|
||||
|
||||
|
||||
class PairedImageDataset(Dataset):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.size = self.get_config('size', 512)
|
||||
self.path = self.get_config('path', required=True)
|
||||
self.default_prompt = self.get_config('default_prompt', '')
|
||||
self.network_weight = self.get_config('network_weight', 1.0)
|
||||
self.file_list = [os.path.join(self.path, file) for file in os.listdir(self.path) if
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
def __len__(self):
|
||||
return len(self.file_list)
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
if key in self.config:
|
||||
value = self.config[key]
|
||||
return value
|
||||
elif required:
|
||||
raise ValueError(f'config file error. Missing "config.dataset.{key}" key')
|
||||
else:
|
||||
return default
|
||||
|
||||
def __getitem__(self, index):
|
||||
img_path = self.file_list[index]
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
|
||||
# see if prompt file exists
|
||||
path_no_ext = os.path.splitext(img_path)[0]
|
||||
prompt_path = path_no_ext + '.txt'
|
||||
if os.path.exists(prompt_path):
|
||||
with open(prompt_path, 'r', encoding='utf-8') as f:
|
||||
prompt = f.read()
|
||||
# remove any newlines
|
||||
prompt = prompt.replace('\n', ', ')
|
||||
# remove new lines for all operating systems
|
||||
prompt = prompt.replace('\r', ', ')
|
||||
prompt_split = prompt.split(',')
|
||||
# remove empty strings
|
||||
prompt_split = [p.strip() for p in prompt_split if p.strip()]
|
||||
# join back together
|
||||
prompt = ', '.join(prompt_split)
|
||||
else:
|
||||
prompt = self.default_prompt
|
||||
|
||||
height = self.size
|
||||
# determine width to keep aspect ratio
|
||||
width = int(img.size[0] * height / img.size[1])
|
||||
|
||||
# Downscale the source image first
|
||||
img = img.resize((width, height), Image.BICUBIC)
|
||||
img = self.transform(img)
|
||||
|
||||
return img, prompt, self.network_weight
|
||||
|
||||
|
||||
51
toolkit/esrgan_utils.py
Normal file
51
toolkit/esrgan_utils.py
Normal file
@@ -0,0 +1,51 @@
|
||||
|
||||
to_basicsr_dict = {
|
||||
'model.0.weight': 'conv_first.weight',
|
||||
'model.0.bias': 'conv_first.bias',
|
||||
'model.1.sub.23.weight': 'conv_body.weight',
|
||||
'model.1.sub.23.bias': 'conv_body.bias',
|
||||
'model.3.weight': 'conv_up1.weight',
|
||||
'model.3.bias': 'conv_up1.bias',
|
||||
'model.6.weight': 'conv_up2.weight',
|
||||
'model.6.bias': 'conv_up2.bias',
|
||||
'model.8.weight': 'conv_hr.weight',
|
||||
'model.8.bias': 'conv_hr.bias',
|
||||
'model.10.bias': 'conv_last.bias',
|
||||
'model.10.weight': 'conv_last.weight',
|
||||
# 'model.1.sub.0.RDB1.conv1.0.weight': 'body.0.rdb1.conv1.weight'
|
||||
}
|
||||
|
||||
def convert_state_dict_to_basicsr(state_dict):
|
||||
new_state_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if k in to_basicsr_dict:
|
||||
new_state_dict[to_basicsr_dict[k]] = v
|
||||
elif k.startswith('model.1.sub.'):
|
||||
bsr_name = k.replace('model.1.sub.', 'body.').lower()
|
||||
bsr_name = bsr_name.replace('.0.weight', '.weight')
|
||||
bsr_name = bsr_name.replace('.0.bias', '.bias')
|
||||
new_state_dict[bsr_name] = v
|
||||
else:
|
||||
new_state_dict[k] = v
|
||||
return new_state_dict
|
||||
|
||||
|
||||
# just matching a commonly used format
|
||||
def convert_basicsr_state_dict_to_save_format(state_dict):
|
||||
new_state_dict = {}
|
||||
to_basicsr_dict_values = list(to_basicsr_dict.values())
|
||||
for k, v in state_dict.items():
|
||||
if k in to_basicsr_dict_values:
|
||||
for key, value in to_basicsr_dict.items():
|
||||
if value == k:
|
||||
new_state_dict[key] = v
|
||||
|
||||
elif k.startswith('body.'):
|
||||
bsr_name = k.replace('body.', 'model.1.sub.').lower()
|
||||
bsr_name = bsr_name.replace('rdb', 'RDB')
|
||||
bsr_name = bsr_name.replace('.weight', '.0.weight')
|
||||
bsr_name = bsr_name.replace('.bias', '.0.bias')
|
||||
new_state_dict[bsr_name] = v
|
||||
else:
|
||||
new_state_dict[k] = v
|
||||
return new_state_dict
|
||||
57
toolkit/extension.py
Normal file
57
toolkit/extension.py
Normal file
@@ -0,0 +1,57 @@
|
||||
import os
|
||||
import importlib
|
||||
import pkgutil
|
||||
from typing import List
|
||||
|
||||
from toolkit.paths import TOOLKIT_ROOT
|
||||
|
||||
|
||||
class Extension(object):
|
||||
"""Base class for extensions.
|
||||
|
||||
Extensions are registered with the ExtensionManager, which is
|
||||
responsible for calling the extension's load() and unload()
|
||||
methods at the appropriate times.
|
||||
|
||||
"""
|
||||
|
||||
name: str = None
|
||||
uid: str = None
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# extend in subclass
|
||||
pass
|
||||
|
||||
|
||||
def get_all_extensions() -> List[Extension]:
|
||||
extension_folders = ['extensions', 'extensions_built_in']
|
||||
|
||||
# This will hold the classes from all extension modules
|
||||
all_extension_classes: List[Extension] = []
|
||||
|
||||
# Iterate over all directories (i.e., packages) in the "extensions" directory
|
||||
for sub_dir in extension_folders:
|
||||
extensions_dir = os.path.join(TOOLKIT_ROOT, sub_dir)
|
||||
for (_, name, _) in pkgutil.iter_modules([extensions_dir]):
|
||||
try:
|
||||
# Import the module
|
||||
module = importlib.import_module(f"{sub_dir}.{name}")
|
||||
# Get the value of the AI_TOOLKIT_EXTENSIONS variable
|
||||
extensions = getattr(module, "AI_TOOLKIT_EXTENSIONS", None)
|
||||
# Check if the value is a list
|
||||
if isinstance(extensions, list):
|
||||
# Iterate over the list and add the classes to the main list
|
||||
all_extension_classes.extend(extensions)
|
||||
except ImportError as e:
|
||||
print(f"Failed to import the {name} module. Error: {str(e)}")
|
||||
|
||||
return all_extension_classes
|
||||
|
||||
|
||||
def get_all_extensions_process_dict():
|
||||
all_extensions = get_all_extensions()
|
||||
process_dict = {}
|
||||
for extension in all_extensions:
|
||||
process_dict[extension.uid] = extension.get_process()
|
||||
return process_dict
|
||||
@@ -1,8 +1,13 @@
|
||||
from typing import Union, OrderedDict
|
||||
|
||||
from toolkit.config import get_config
|
||||
|
||||
|
||||
def get_job(config_path):
|
||||
config = get_config(config_path)
|
||||
def get_job(
|
||||
config_path: Union[str, dict, OrderedDict],
|
||||
name=None
|
||||
):
|
||||
config = get_config(config_path, name)
|
||||
if not config['job']:
|
||||
raise ValueError('config file is invalid. Missing "job" key')
|
||||
|
||||
@@ -13,9 +18,27 @@ def get_job(config_path):
|
||||
if job == 'train':
|
||||
from jobs import TrainJob
|
||||
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':
|
||||
# from jobs import TrainJob
|
||||
# return TrainJob(config)
|
||||
else:
|
||||
raise ValueError(f'Unknown job type {job}')
|
||||
|
||||
|
||||
def run_job(
|
||||
config: Union[str, dict, OrderedDict],
|
||||
name=None
|
||||
):
|
||||
job = get_job(config, name)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
|
||||
@@ -892,6 +892,9 @@ def convert_ldm_clip_checkpoint_v1(checkpoint):
|
||||
for key in keys:
|
||||
if key.startswith("cond_stage_model.transformer"):
|
||||
text_model_dict[key[len("cond_stage_model.transformer."):]] = checkpoint[key]
|
||||
# support checkpoint without position_ids (invalid checkpoint)
|
||||
if "text_model.embeddings.position_ids" not in text_model_dict:
|
||||
text_model_dict["text_model.embeddings.position_ids"] = torch.arange(77).unsqueeze(0) # 77 is the max length of the text
|
||||
return text_model_dict
|
||||
|
||||
|
||||
@@ -1257,6 +1260,10 @@ def load_models_from_stable_diffusion_checkpoint(v2, ckpt_path, device="cpu", dt
|
||||
text_model = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14").to(device)
|
||||
logging.set_verbosity_warning()
|
||||
|
||||
# latest transformers doesnt have position ids. Do we remove it?
|
||||
if "text_model.embeddings.position_ids" not in text_model.state_dict():
|
||||
del converted_text_encoder_checkpoint["text_model.embeddings.position_ids"]
|
||||
|
||||
info = text_model.load_state_dict(converted_text_encoder_checkpoint)
|
||||
print("loading text encoder:", info)
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
class ReductionKernel(nn.Module):
|
||||
@@ -29,3 +30,15 @@ class ReductionKernel(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return nn.functional.conv2d(x, self.kernel, stride=self.kernel_size, padding=0, groups=1)
|
||||
|
||||
|
||||
class CheckpointGradients(nn.Module):
|
||||
def __init__(self, is_gradient_checkpointing=True):
|
||||
super(CheckpointGradients, self).__init__()
|
||||
self.is_gradient_checkpointing = is_gradient_checkpointing
|
||||
|
||||
def forward(self, module, *args, num_chunks=1):
|
||||
if self.is_gradient_checkpointing:
|
||||
return checkpoint(module, *args, num_chunks=self.num_chunks)
|
||||
else:
|
||||
return module(*args)
|
||||
|
||||
@@ -6,12 +6,14 @@
|
||||
import os
|
||||
import math
|
||||
from typing import Optional, List, Type, Set, Literal
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers import UNet2DConditionModel
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from toolkit.metadata import add_model_hash_to_meta
|
||||
|
||||
UNET_TARGET_REPLACE_MODULE_TRANSFORMER = [
|
||||
"Transformer2DModel", # どうやらこっちの方らしい? # attn1, 2
|
||||
@@ -31,7 +33,7 @@ TRAINING_METHODS = Literal[
|
||||
"innoxattn", # train all layers except self attention layers
|
||||
"selfattn", # ESD-u, train only self attention layers
|
||||
"xattn", # ESD-x, train only x attention layers
|
||||
"full", # train all layers
|
||||
"full", # train all layers
|
||||
# "notime",
|
||||
# "xlayer",
|
||||
# "outxattn",
|
||||
@@ -48,12 +50,12 @@ class LoRAModule(nn.Module):
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
lora_name,
|
||||
org_module: nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
self,
|
||||
lora_name,
|
||||
org_module: nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
):
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
super().__init__()
|
||||
@@ -102,19 +104,19 @@ class LoRAModule(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return (
|
||||
self.org_forward(x)
|
||||
+ self.lora_up(self.lora_down(x)) * self.multiplier * self.scale
|
||||
self.org_forward(x)
|
||||
+ self.lora_up(self.lora_down(x)) * self.multiplier * self.scale
|
||||
)
|
||||
|
||||
|
||||
class LoRANetwork(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
unet: UNet2DConditionModel,
|
||||
rank: int = 4,
|
||||
multiplier: float = 1.0,
|
||||
alpha: float = 1.0,
|
||||
train_method: TRAINING_METHODS = "full",
|
||||
self,
|
||||
unet: UNet2DConditionModel,
|
||||
rank: int = 4,
|
||||
multiplier: float = 1.0,
|
||||
alpha: float = 1.0,
|
||||
train_method: TRAINING_METHODS = "full",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -140,7 +142,7 @@ class LoRANetwork(nn.Module):
|
||||
lora_names = set()
|
||||
for lora in self.unet_loras:
|
||||
assert (
|
||||
lora.lora_name not in lora_names
|
||||
lora.lora_name not in lora_names
|
||||
), f"duplicated lora name: {lora.lora_name}. {lora_names}"
|
||||
lora_names.add(lora.lora_name)
|
||||
|
||||
@@ -157,13 +159,13 @@ class LoRANetwork(nn.Module):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def create_modules(
|
||||
self,
|
||||
prefix: str,
|
||||
root_module: nn.Module,
|
||||
target_replace_modules: List[str],
|
||||
rank: int,
|
||||
multiplier: float,
|
||||
train_method: TRAINING_METHODS,
|
||||
self,
|
||||
prefix: str,
|
||||
root_module: nn.Module,
|
||||
target_replace_modules: List[str],
|
||||
rank: int,
|
||||
multiplier: float,
|
||||
train_method: TRAINING_METHODS,
|
||||
) -> list:
|
||||
loras = []
|
||||
|
||||
@@ -212,6 +214,8 @@ class LoRANetwork(nn.Module):
|
||||
|
||||
def save_weights(self, file, dtype=None, metadata: Optional[dict] = None):
|
||||
state_dict = self.state_dict()
|
||||
if metadata is None:
|
||||
metadata = OrderedDict()
|
||||
|
||||
if dtype is not None:
|
||||
for key in list(state_dict.keys()):
|
||||
@@ -221,9 +225,10 @@ class LoRANetwork(nn.Module):
|
||||
|
||||
for key in list(state_dict.keys()):
|
||||
if not key.startswith("lora"):
|
||||
# lora以外除外
|
||||
# remove any not lora
|
||||
del state_dict[key]
|
||||
|
||||
metadata = add_model_hash_to_meta(state_dict, metadata)
|
||||
if os.path.splitext(file)[1] == ".safetensors":
|
||||
save_file(state_dict, file, metadata)
|
||||
else:
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from typing import List, Optional, Dict, Type, Union
|
||||
|
||||
@@ -9,7 +11,174 @@ from .paths import SD_SCRIPTS_ROOT
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from networks.lora import LoRANetwork, LoRAModule, get_block_index
|
||||
from networks.lora import LoRANetwork, get_block_index
|
||||
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
|
||||
|
||||
|
||||
class LoRAModule(torch.nn.Module):
|
||||
"""
|
||||
replaces forward method of the original Linear, instead of replacing the original Linear module.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
lora_name,
|
||||
org_module: torch.nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
dropout=None,
|
||||
rank_dropout=None,
|
||||
module_dropout=None,
|
||||
):
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
super().__init__()
|
||||
self.lora_name = lora_name
|
||||
|
||||
if org_module.__class__.__name__ == "Conv2d":
|
||||
in_dim = org_module.in_channels
|
||||
out_dim = org_module.out_channels
|
||||
else:
|
||||
in_dim = org_module.in_features
|
||||
out_dim = org_module.out_features
|
||||
|
||||
# if limit_rank:
|
||||
# self.lora_dim = min(lora_dim, in_dim, out_dim)
|
||||
# if self.lora_dim != lora_dim:
|
||||
# print(f"{lora_name} dim (rank) is changed to: {self.lora_dim}")
|
||||
# else:
|
||||
self.lora_dim = lora_dim
|
||||
|
||||
if org_module.__class__.__name__ == "Conv2d":
|
||||
kernel_size = org_module.kernel_size
|
||||
stride = org_module.stride
|
||||
padding = org_module.padding
|
||||
self.lora_down = torch.nn.Conv2d(in_dim, self.lora_dim, kernel_size, stride, padding, bias=False)
|
||||
self.lora_up = torch.nn.Conv2d(self.lora_dim, out_dim, (1, 1), (1, 1), bias=False)
|
||||
else:
|
||||
self.lora_down = torch.nn.Linear(in_dim, self.lora_dim, bias=False)
|
||||
self.lora_up = torch.nn.Linear(self.lora_dim, out_dim, bias=False)
|
||||
|
||||
if type(alpha) == torch.Tensor:
|
||||
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
|
||||
self.scale = alpha / self.lora_dim
|
||||
self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える
|
||||
|
||||
# same as microsoft's
|
||||
torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5))
|
||||
torch.nn.init.zeros_(self.lora_up.weight)
|
||||
|
||||
self.multiplier: Union[float, List[float]] = multiplier
|
||||
self.org_module = org_module # remove in applying
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
def apply_to(self):
|
||||
self.org_forward = self.org_module.forward
|
||||
self.org_module.forward = self.forward
|
||||
del self.org_module
|
||||
|
||||
# this allows us to set different multipliers on a per item in a batch basis
|
||||
# allowing us to run positive and negative weights in the same batch
|
||||
# really only useful for slider training for now
|
||||
def get_multiplier(self, lora_up):
|
||||
batch_size = lora_up.size(0)
|
||||
# batch will have all negative prompts first and positive prompts second
|
||||
# our multiplier list is for a prompt pair. So we need to repeat it for positive and negative prompts
|
||||
# if there is more than our multiplier, it is liekly a batch size increase, so we need to
|
||||
# interleve the multipliers
|
||||
if isinstance(self.multiplier, list):
|
||||
if len(self.multiplier) == 0:
|
||||
# single item, just return it
|
||||
return self.multiplier[0]
|
||||
elif len(self.multiplier) == batch_size:
|
||||
# not doing CFG
|
||||
multiplier_tensor = torch.tensor(self.multiplier).to(lora_up.device, dtype=lora_up.dtype)
|
||||
else:
|
||||
|
||||
# we have a list of multipliers, so we need to get the multiplier for this batch
|
||||
multiplier_tensor = torch.tensor(self.multiplier * 2).to(lora_up.device, dtype=lora_up.dtype)
|
||||
# should be 1 for if total batch size was 1
|
||||
num_interleaves = (batch_size // 2) // len(self.multiplier)
|
||||
multiplier_tensor = multiplier_tensor.repeat_interleave(num_interleaves)
|
||||
|
||||
# match lora_up rank
|
||||
if len(lora_up.size()) == 2:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1)
|
||||
elif len(lora_up.size()) == 3:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1, 1)
|
||||
elif len(lora_up.size()) == 4:
|
||||
multiplier_tensor = multiplier_tensor.view(-1, 1, 1, 1)
|
||||
return multiplier_tensor
|
||||
|
||||
else:
|
||||
return self.multiplier
|
||||
|
||||
def _call_forward(self, x):
|
||||
# module dropout
|
||||
if self.module_dropout is not None and self.training:
|
||||
if torch.rand(1) < self.module_dropout:
|
||||
return 0.0 # added to original forward
|
||||
|
||||
lx = self.lora_down(x)
|
||||
|
||||
# normal dropout
|
||||
if self.dropout is not None and self.training:
|
||||
lx = torch.nn.functional.dropout(lx, p=self.dropout)
|
||||
|
||||
# rank dropout
|
||||
if self.rank_dropout is not None and self.training:
|
||||
mask = torch.rand((lx.size(0), self.lora_dim), device=lx.device) > self.rank_dropout
|
||||
if len(lx.size()) == 3:
|
||||
mask = mask.unsqueeze(1) # for Text Encoder
|
||||
elif len(lx.size()) == 4:
|
||||
mask = mask.unsqueeze(-1).unsqueeze(-1) # for Conv2d
|
||||
lx = lx * mask
|
||||
|
||||
# scaling for rank dropout: treat as if the rank is changed
|
||||
# maskから計算することも考えられるが、augmentation的な効果を期待してrank_dropoutを用いる
|
||||
scale = self.scale * (1.0 / (1.0 - self.rank_dropout)) # redundant for readability
|
||||
else:
|
||||
scale = self.scale
|
||||
|
||||
lx = self.lora_up(lx)
|
||||
|
||||
multiplier = self.get_multiplier(lx)
|
||||
|
||||
return lx * multiplier * scale
|
||||
|
||||
def create_custom_forward(self):
|
||||
def custom_forward(*inputs):
|
||||
return self._call_forward(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
def forward(self, x):
|
||||
org_forwarded = self.org_forward(x)
|
||||
# TODO this just loses the grad. Not sure why. Probably why no one else is doing it either
|
||||
# if torch.is_grad_enabled() and self.is_checkpointing and self.training:
|
||||
# lora_output = checkpoint(
|
||||
# self.create_custom_forward(),
|
||||
# x,
|
||||
# )
|
||||
# else:
|
||||
# lora_output = self._call_forward(x)
|
||||
|
||||
lora_output = self._call_forward(x)
|
||||
|
||||
return org_forwarded + lora_output
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.is_checkpointing = True
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
self.is_checkpointing = False
|
||||
|
||||
|
||||
class LoRASpecialNetwork(LoRANetwork):
|
||||
@@ -70,6 +239,7 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
if modules_dim is not None:
|
||||
print(f"create LoRA network from weights")
|
||||
@@ -236,11 +406,11 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
torch.save(state_dict, file)
|
||||
|
||||
@property
|
||||
def multiplier(self):
|
||||
def multiplier(self) -> Union[float, List[float]]:
|
||||
return self._multiplier
|
||||
|
||||
@multiplier.setter
|
||||
def multiplier(self, value):
|
||||
def multiplier(self, value: Union[float, List[float]]):
|
||||
self._multiplier = value
|
||||
self._update_lora_multiplier()
|
||||
|
||||
@@ -261,6 +431,8 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.multiplier = 0
|
||||
|
||||
# called when the context manager is entered
|
||||
# ie: with network:
|
||||
def __enter__(self):
|
||||
self.is_active = True
|
||||
self._update_lora_multiplier()
|
||||
@@ -278,3 +450,29 @@ class LoRASpecialNetwork(LoRANetwork):
|
||||
loras += self.text_encoder_loras
|
||||
for lora in loras:
|
||||
lora.to(device, dtype)
|
||||
|
||||
def _update_checkpointing(self):
|
||||
if self.is_checkpointing:
|
||||
if hasattr(self, 'unet_loras'):
|
||||
for lora in self.unet_loras:
|
||||
lora.enable_gradient_checkpointing()
|
||||
if hasattr(self, 'text_encoder_loras'):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.enable_gradient_checkpointing()
|
||||
else:
|
||||
if hasattr(self, 'unet_loras'):
|
||||
for lora in self.unet_loras:
|
||||
lora.disable_gradient_checkpointing()
|
||||
if hasattr(self, 'text_encoder_loras'):
|
||||
for lora in self.text_encoder_loras:
|
||||
lora.disable_gradient_checkpointing()
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
# not supported
|
||||
self.is_checkpointing = True
|
||||
self._update_checkpointing()
|
||||
|
||||
def disable_gradient_checkpointing(self):
|
||||
# not supported
|
||||
self.is_checkpointing = False
|
||||
self._update_checkpointing()
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -1,18 +1,23 @@
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
from io import BytesIO
|
||||
|
||||
import safetensors
|
||||
from safetensors import safe_open
|
||||
|
||||
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
|
||||
meta_string = json.dumps(meta)
|
||||
if name is not None:
|
||||
meta_string = meta_string.replace("[name]", name)
|
||||
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
|
||||
for key, value in save_meta.items():
|
||||
# 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
|
||||
|
||||
|
||||
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:
|
||||
parsed_meta = OrderedDict()
|
||||
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)
|
||||
# warn user that dadaptation is deprecated
|
||||
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"):
|
||||
import bitsandbytes
|
||||
|
||||
@@ -43,6 +54,8 @@ def get_optimizer(
|
||||
elif lower_type == 'lion':
|
||||
from lion_pytorch import Lion
|
||||
return Lion(params, lr=learning_rate, **optimizer_params)
|
||||
elif lower_type == 'adagrad':
|
||||
optimizer = torch.optim.Adagrad(params, lr=float(learning_rate), **optimizer_params)
|
||||
else:
|
||||
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
||||
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")
|
||||
REPOS_ROOT = os.path.join(TOOLKIT_ROOT, "repositories")
|
||||
|
||||
# check if ENV variable is set
|
||||
if 'MODELS_PATH' in os.environ:
|
||||
MODELS_PATH = os.environ['MODELS_PATH']
|
||||
else:
|
||||
MODELS_PATH = os.path.join(TOOLKIT_ROOT, "models")
|
||||
|
||||
|
||||
def get_path(path):
|
||||
# we allow absolute paths, but if it is not absolute, we assume it is relative to the toolkit root
|
||||
|
||||
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 os
|
||||
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||
from safetensors.torch import save_file
|
||||
from torch.amp import autocast
|
||||
from tqdm import tqdm
|
||||
from torchvision.transforms import Resize
|
||||
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from library.sdxl_train_util import _load_target_model as load_sdxl_target_model
|
||||
from toolkit import train_tools
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.model_util_sdxl import load_models_from_sdxl_checkpoint
|
||||
from toolkit.paths import REPOS_ROOT, MODELS_PATH
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'sd-scripts'))
|
||||
from leco import train_util
|
||||
import torch
|
||||
from library import model_util
|
||||
from library.sdxl_model_util import convert_text_encoder_2_state_dict_to_sdxl
|
||||
from library.model_util import convert_unet_state_dict_to_sd, convert_text_encoder_state_dict_to_sd_v2
|
||||
from library import model_util, train_util as kohya_train_util, sdxl_train_util
|
||||
from library.sdxl_model_util import convert_text_encoder_2_state_dict_to_sdxl, DIFFUSERS_SDXL_UNET_CONFIG
|
||||
from diffusers.schedulers import DDPMScheduler
|
||||
from toolkit.pipelines import CustomStableDiffusionXLPipeline, CustomStableDiffusionPipeline, \
|
||||
HackedStableDiffusionXLPipeline
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
import diffusers
|
||||
|
||||
# tell it to shut up
|
||||
diffusers.logging.set_verbosity(diffusers.logging.ERROR)
|
||||
|
||||
|
||||
class BlankNetwork:
|
||||
multiplier = 1.0
|
||||
is_active = True
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __enter__(self):
|
||||
self.is_active = True
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.is_active = False
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
UNET_IN_CHANNELS = 4 # Stable Diffusion の in_channels は 4 で固定。XLも同じ。
|
||||
VAE_SCALE_FACTOR = 8 # 2 ** (len(vae.config.block_out_channels) - 1) = 8
|
||||
|
||||
|
||||
class PromptEmbeds:
|
||||
text_embeds: torch.FloatTensor
|
||||
pooled_embeds: Union[torch.FloatTensor, None]
|
||||
text_embeds: torch.Tensor
|
||||
pooled_embeds: Union[torch.Tensor, None]
|
||||
|
||||
def __init__(self, args) -> None:
|
||||
def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Tensor]) -> None:
|
||||
if isinstance(args, list) or isinstance(args, tuple):
|
||||
# xl
|
||||
self.text_embeds = args[0]
|
||||
@@ -29,30 +73,536 @@ class PromptEmbeds:
|
||||
self.text_embeds = args
|
||||
self.pooled_embeds = None
|
||||
|
||||
def to(self, **kwargs):
|
||||
self.text_embeds = self.text_embeds.to(**kwargs)
|
||||
def to(self, *args, **kwargs):
|
||||
self.text_embeds = self.text_embeds.to(*args, **kwargs)
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
pipeline: Union[None, 'StableDiffusionPipeline', 'CustomStableDiffusionXLPipeline']
|
||||
vae: Union[None, 'AutoencoderKL']
|
||||
unet: Union[None, 'UNet2DConditionModel']
|
||||
text_encoder: Union[None, 'CLIPTextModel', List[Union['CLIPTextModel', 'CLIPTextModelWithProjection']]]
|
||||
tokenizer: Union[None, 'CLIPTokenizer', List['CLIPTokenizer']]
|
||||
noise_scheduler: Union[None, 'KarrasDiffusionSchedulers', 'DDPMScheduler']
|
||||
device: str
|
||||
dtype: str
|
||||
torch_dtype: torch.dtype
|
||||
device_torch: torch.device
|
||||
model_config: ModelConfig
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
unet,
|
||||
noise_scheduler,
|
||||
is_xl=False
|
||||
device,
|
||||
model_config: ModelConfig,
|
||||
dtype='fp16',
|
||||
custom_pipeline=None,
|
||||
):
|
||||
# text encoder has a list of 2 for xl
|
||||
self.vae = vae
|
||||
self.tokenizer = tokenizer
|
||||
self.text_encoder = text_encoder
|
||||
self.unet = unet
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.is_xl = is_xl
|
||||
self.custom_pipeline = custom_pipeline
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.torch_dtype = get_torch_dtype(dtype)
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.model_config = model_config
|
||||
self.prediction_type = "v_prediction" if self.model_config.is_v_pred else "epsilon"
|
||||
|
||||
# sdxl stuff
|
||||
self.logit_scale = None
|
||||
self.ckpt_info = None
|
||||
self.is_loaded = False
|
||||
|
||||
# to hold network if there is one
|
||||
self.network = None
|
||||
self.is_xl = model_config.is_xl
|
||||
self.is_v2 = model_config.is_v2
|
||||
|
||||
def load_model(self):
|
||||
if self.is_loaded:
|
||||
return
|
||||
dtype = get_torch_dtype(self.dtype)
|
||||
|
||||
# TODO handle other schedulers
|
||||
# sch = KDPM2DiscreteScheduler
|
||||
sch = DDPMScheduler
|
||||
# do our own scheduler
|
||||
prediction_type = "v_prediction" if self.model_config.is_v_pred else "epsilon"
|
||||
scheduler = sch(
|
||||
num_train_timesteps=1000,
|
||||
beta_start=0.00085,
|
||||
beta_end=0.0120,
|
||||
beta_schedule="scaled_linear",
|
||||
clip_sample=False,
|
||||
prediction_type=prediction_type,
|
||||
steps_offset=1
|
||||
)
|
||||
|
||||
model_path = self.model_config.name_or_path
|
||||
if 'civitai.com' in self.model_config.name_or_path:
|
||||
# load is a civit ai model, use the loader.
|
||||
from toolkit.civitai import get_model_path_from_url
|
||||
model_path = get_model_path_from_url(self.model_config.name_or_path)
|
||||
|
||||
if self.model_config.is_xl:
|
||||
# load from kohya
|
||||
# (
|
||||
# load_stable_diffusion_format,
|
||||
# text_encoder1,
|
||||
# text_encoder2,
|
||||
# vae,
|
||||
# unet,
|
||||
# logit_scale,
|
||||
# ckpt_info,
|
||||
# ) = load_sdxl_target_model(
|
||||
# model_path,
|
||||
# self.model_config.vae_path if self.model_config.vae_path is not None else model_path,
|
||||
# 'sdxl',
|
||||
# self.dtype,
|
||||
# self.device,
|
||||
# )
|
||||
|
||||
(
|
||||
text_encoder1,
|
||||
text_encoder2,
|
||||
vae,
|
||||
unet,
|
||||
logit_scale,
|
||||
ckpt_info,
|
||||
) = load_models_from_sdxl_checkpoint(
|
||||
'sdxl',
|
||||
model_path,
|
||||
self.device
|
||||
)
|
||||
|
||||
class Config:
|
||||
def __init__(self):
|
||||
# add all items from DIFFUSERS_SDXL_UNET_CONFIG as attributes
|
||||
for k, v in DIFFUSERS_SDXL_UNET_CONFIG.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
# add diffusers stuff
|
||||
unet.config = Config()
|
||||
|
||||
if self.model_config.vae_path is not None:
|
||||
vae = model_util.load_vae(self.model_config.vae_path, self.dtype)
|
||||
print("additional VAE loaded")
|
||||
|
||||
text_encoder1, text_encoder2, unet = kohya_train_util.transform_models_if_DDP(
|
||||
[text_encoder1, text_encoder2, unet]
|
||||
)
|
||||
class Args:
|
||||
tokenizer_cache_dir = os.path.join(MODELS_PATH, 'clip_cache')
|
||||
|
||||
args = Args()
|
||||
os.makedirs(args.tokenizer_cache_dir, exist_ok=True)
|
||||
|
||||
self.tokenizer = sdxl_train_util.load_tokenizers(args)
|
||||
# tokenizer1 = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
|
||||
# tokenizer2 = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer_2")
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
self.vae = vae
|
||||
self.unet = unet.to(self.device_torch, dtype=dtype)
|
||||
self.text_encoder = [text_encoder1, text_encoder2]
|
||||
self.logit_scale = logit_scale
|
||||
self.ckpt_info = ckpt_info
|
||||
|
||||
|
||||
|
||||
# if self.custom_pipeline is not None:
|
||||
# pipln = self.custom_pipeline
|
||||
# else:
|
||||
# pipln = CustomStableDiffusionXLPipeline
|
||||
#
|
||||
# # see if path exists
|
||||
# if not os.path.exists(model_path):
|
||||
# # try to load with default diffusers
|
||||
# pipe = pipln.from_pretrained(
|
||||
# model_path,
|
||||
# dtype=dtype,
|
||||
# scheduler_type='ddpm',
|
||||
# device=self.device_torch,
|
||||
# ).to(self.device_torch)
|
||||
# else:
|
||||
# pipe = pipln.from_single_file(
|
||||
# model_path,
|
||||
# dtype=dtype,
|
||||
# scheduler_type='ddpm',
|
||||
# device=self.device_torch,
|
||||
# ).to(self.device_torch)
|
||||
#
|
||||
# text_encoders = [pipe.text_encoder, pipe.text_encoder_2]
|
||||
# tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
|
||||
# for text_encoder in text_encoders:
|
||||
# text_encoder.to(self.device_torch, dtype=dtype)
|
||||
# text_encoder.requires_grad_(False)
|
||||
# text_encoder.eval()
|
||||
# text_encoder = text_encoders
|
||||
else:
|
||||
if self.custom_pipeline is not None:
|
||||
pipln = self.custom_pipeline
|
||||
else:
|
||||
pipln = CustomStableDiffusionPipeline
|
||||
|
||||
# see if path exists
|
||||
if not os.path.exists(model_path):
|
||||
# try to load with default diffusers
|
||||
pipe = pipln.from_pretrained(
|
||||
model_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='dpm',
|
||||
device=self.device_torch,
|
||||
load_safety_checker=False,
|
||||
requires_safety_checker=False,
|
||||
safety_checker=False
|
||||
).to(self.device_torch)
|
||||
else:
|
||||
pipe = pipln.from_single_file(
|
||||
model_path,
|
||||
dtype=dtype,
|
||||
scheduler_type='dpm',
|
||||
device=self.device_torch,
|
||||
load_safety_checker=False,
|
||||
requires_safety_checker=False,
|
||||
safety_checker=False
|
||||
).to(self.device_torch)
|
||||
|
||||
pipe.register_to_config(requires_safety_checker=False)
|
||||
text_encoder = pipe.text_encoder
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
tokenizer = pipe.tokenizer
|
||||
self.tokenizer = tokenizer
|
||||
self.text_encoder = text_encoder
|
||||
self.pipeline = pipe
|
||||
|
||||
# scheduler doesn't get set sometimes, so we set it here
|
||||
pipe.scheduler = scheduler
|
||||
|
||||
self.unet = pipe.unet
|
||||
self.vae = pipe.vae.to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.noise_scheduler = scheduler
|
||||
self.vae.eval()
|
||||
self.vae.requires_grad_(False)
|
||||
self.unet.to(self.device_torch, dtype=dtype)
|
||||
self.unet.requires_grad_(False)
|
||||
self.unet.eval()
|
||||
self.is_loaded = True
|
||||
|
||||
def generate_images(self, image_configs: List[GenerateImageConfig]):
|
||||
# sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if self.network is not None:
|
||||
self.network.eval()
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
# save current seed state for training
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
original_device_dict = {
|
||||
'vae': self.vae.device,
|
||||
'unet': self.unet.device,
|
||||
# 'tokenizer': self.tokenizer.device,
|
||||
}
|
||||
|
||||
# handle sdxl text encoder
|
||||
if isinstance(self.text_encoder, list):
|
||||
for encoder, i in zip(self.text_encoder, range(len(self.text_encoder))):
|
||||
original_device_dict[f'text_encoder_{i}'] = encoder.device
|
||||
encoder.to(self.device_torch)
|
||||
else:
|
||||
original_device_dict['text_encoder'] = self.text_encoder.device
|
||||
self.text_encoder.to(self.device_torch)
|
||||
|
||||
self.vae.to(self.device_torch)
|
||||
self.unet.to(self.device_torch)
|
||||
|
||||
# TODO add clip skip
|
||||
if self.is_xl:
|
||||
pipeline = HackedStableDiffusionXLPipeline(
|
||||
vae=self.vae,
|
||||
unet=self.unet,
|
||||
text_encoder=self.text_encoder[0],
|
||||
text_encoder_2=self.text_encoder[1],
|
||||
tokenizer=self.tokenizer[0],
|
||||
tokenizer_2=self.tokenizer[1],
|
||||
scheduler=self.noise_scheduler,
|
||||
add_watermarker=False,
|
||||
).to(self.device_torch)
|
||||
# force turn that (ruin your images with obvious green and red dots) the #$@@ off!!!
|
||||
pipeline.watermark = None
|
||||
else:
|
||||
pipeline = StableDiffusionPipeline(
|
||||
vae=self.vae,
|
||||
unet=self.unet,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=self.noise_scheduler,
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
).to(self.device_torch)
|
||||
# disable progress bar
|
||||
pipeline.set_progress_bar_config(disable=True)
|
||||
|
||||
start_multiplier = 1.0
|
||||
if self.network is not None:
|
||||
start_multiplier = self.network.multiplier
|
||||
|
||||
pipeline.to(self.device_torch)
|
||||
with network:
|
||||
with torch.no_grad():
|
||||
if self.network is not None:
|
||||
assert self.network.is_active
|
||||
|
||||
for i in tqdm(range(len(image_configs)), desc=f"Generating Images", leave=False):
|
||||
gen_config = image_configs[i]
|
||||
|
||||
if self.network is not None:
|
||||
self.network.multiplier = gen_config.network_multiplier
|
||||
torch.manual_seed(gen_config.seed)
|
||||
torch.cuda.manual_seed(gen_config.seed)
|
||||
|
||||
if self.is_xl:
|
||||
with autocast('cuda'):
|
||||
img = pipeline(
|
||||
prompt=gen_config.prompt,
|
||||
prompt_2=gen_config.prompt_2,
|
||||
negative_prompt=gen_config.negative_prompt,
|
||||
negative_prompt_2=gen_config.negative_prompt_2,
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
guidance_rescale=gen_config.guidance_rescale,
|
||||
).images[0]
|
||||
else:
|
||||
img = pipeline(
|
||||
prompt=gen_config.prompt,
|
||||
negative_prompt=gen_config.negative_prompt,
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
).images[0]
|
||||
|
||||
gen_config.save_image(img)
|
||||
|
||||
# clear pipeline and cache to reduce vram usage
|
||||
del pipeline
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# restore training state
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
|
||||
self.vae.to(original_device_dict['vae'])
|
||||
self.unet.to(original_device_dict['unet'])
|
||||
if isinstance(self.text_encoder, list):
|
||||
for encoder, i in zip(self.text_encoder, range(len(self.text_encoder))):
|
||||
encoder.to(original_device_dict[f'text_encoder_{i}'])
|
||||
else:
|
||||
self.text_encoder.to(original_device_dict['text_encoder'])
|
||||
if self.network is not None:
|
||||
self.network.train()
|
||||
self.network.multiplier = start_multiplier
|
||||
# self.tokenizer.to(original_device_dict['tokenizer'])
|
||||
|
||||
def get_latent_noise(
|
||||
self,
|
||||
height=None,
|
||||
width=None,
|
||||
pixel_height=None,
|
||||
pixel_width=None,
|
||||
batch_size=1,
|
||||
noise_offset=0.0,
|
||||
):
|
||||
if height is None and pixel_height is None:
|
||||
raise ValueError("height or pixel_height must be specified")
|
||||
if width is None and pixel_width is None:
|
||||
raise ValueError("width or pixel_width must be specified")
|
||||
if height is None:
|
||||
height = pixel_height // VAE_SCALE_FACTOR
|
||||
if width is None:
|
||||
width = pixel_width // VAE_SCALE_FACTOR
|
||||
|
||||
noise = torch.randn(
|
||||
(
|
||||
batch_size,
|
||||
UNET_IN_CHANNELS,
|
||||
height,
|
||||
width,
|
||||
),
|
||||
device="cpu",
|
||||
)
|
||||
noise = apply_noise_offset(noise, noise_offset)
|
||||
return noise
|
||||
|
||||
def get_time_ids_from_latents(self, latents: torch.Tensor):
|
||||
bs, ch, h, w = list(latents.shape)
|
||||
|
||||
height = h * VAE_SCALE_FACTOR
|
||||
width = w * VAE_SCALE_FACTOR
|
||||
|
||||
dtype = latents.dtype
|
||||
|
||||
if self.is_xl:
|
||||
prompt_ids = train_util.get_add_time_ids(
|
||||
height,
|
||||
width,
|
||||
dynamic_crops=False, # look into this
|
||||
dtype=dtype,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
return prompt_ids
|
||||
else:
|
||||
return None
|
||||
|
||||
def predict_noise(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
text_embeddings: Union[PromptEmbeds, None] = None,
|
||||
timestep: Union[int, torch.Tensor] = 1,
|
||||
guidance_scale=7.5,
|
||||
guidance_rescale=0, # 0.7 sdxl
|
||||
add_time_ids=None,
|
||||
conditional_embeddings: Union[PromptEmbeds, None] = None,
|
||||
unconditional_embeddings: Union[PromptEmbeds, None] = None,
|
||||
**kwargs,
|
||||
):
|
||||
# get the embeddings
|
||||
if text_embeddings is None and conditional_embeddings is None:
|
||||
raise ValueError("Either text_embeddings or conditional_embeddings must be specified")
|
||||
if text_embeddings is None and unconditional_embeddings is not None:
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
unconditional_embeddings, # negative embedding
|
||||
conditional_embeddings, # positive embedding
|
||||
latents.shape[0], # batch size
|
||||
)
|
||||
elif text_embeddings is None and conditional_embeddings is not None:
|
||||
# not doing cfg
|
||||
text_embeddings = conditional_embeddings
|
||||
|
||||
# CFG is comparing neg and positive, if we have concatenated embeddings
|
||||
# then we are doing it, otherwise we are not and takes half the time.
|
||||
do_classifier_free_guidance = True
|
||||
|
||||
# check if batch size of embeddings matches batch size of latents
|
||||
if latents.shape[0] == text_embeddings.text_embeds.shape[0]:
|
||||
do_classifier_free_guidance = False
|
||||
elif latents.shape[0] * 2 != text_embeddings.text_embeds.shape[0]:
|
||||
raise ValueError("Batch size of latents must be the same or half the batch size of text embeddings")
|
||||
|
||||
if self.is_xl:
|
||||
if add_time_ids is None:
|
||||
add_time_ids = self.get_time_ids_from_latents(latents)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# todo check this with larget batches
|
||||
train_util.concat_embeddings(
|
||||
add_time_ids, add_time_ids, 1
|
||||
)
|
||||
else:
|
||||
# concat to fit batch size
|
||||
add_time_ids = torch.cat([add_time_ids] * latents.shape[0])
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
|
||||
latent_model_input = self.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
added_cond_kwargs = {
|
||||
"text_embeds": text_embeddings.pooled_embeds,
|
||||
"time_ids": add_time_ids,
|
||||
}
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
).sample
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# perform guidance
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
# https://github.com/huggingface/diffusers/blob/7a91ea6c2b53f94da930a61ed571364022b21044/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py#L775
|
||||
if guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
|
||||
|
||||
else:
|
||||
if do_classifier_free_guidance:
|
||||
# if we are doing classifier free guidance, need to double up
|
||||
latent_model_input = torch.cat([latents] * 2)
|
||||
else:
|
||||
latent_model_input = latents
|
||||
|
||||
latent_model_input = self.noise_scheduler.scale_model_input(latent_model_input, timestep)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
).sample
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
# perform guidance
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
return noise_pred
|
||||
|
||||
# ref: https://github.com/huggingface/diffusers/blob/0bab447670f47c28df60fbd2f6a0f833f75a16f5/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L746
|
||||
def diffuse_some_steps(
|
||||
self,
|
||||
latents: torch.FloatTensor,
|
||||
text_embeddings: PromptEmbeds,
|
||||
total_timesteps: int = 1000,
|
||||
start_timesteps=0,
|
||||
guidance_scale=1,
|
||||
add_time_ids=None,
|
||||
**kwargs,
|
||||
):
|
||||
|
||||
for timestep in tqdm(self.noise_scheduler.timesteps[start_timesteps:total_timesteps], leave=False):
|
||||
noise_pred = self.predict_noise(
|
||||
latents,
|
||||
text_embeddings,
|
||||
timestep,
|
||||
guidance_scale=guidance_scale,
|
||||
add_time_ids=add_time_ids,
|
||||
**kwargs,
|
||||
)
|
||||
latents = self.noise_scheduler.step(noise_pred, timestep, latents).prev_sample
|
||||
|
||||
# return latents_steps
|
||||
return latents
|
||||
|
||||
def encode_prompt(self, prompt, num_images_per_prompt=1) -> PromptEmbeds:
|
||||
prompt = prompt
|
||||
@@ -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):
|
||||
state_dict = {}
|
||||
|
||||
def update_sd(prefix, sd):
|
||||
for k, v in sd.items():
|
||||
key = prefix + k
|
||||
v = v.detach().clone()
|
||||
state_dict[key] = v.to("cpu", dtype=get_torch_dtype(save_dtype))
|
||||
|
||||
# todo see what logit scale is
|
||||
if self.is_xl:
|
||||
|
||||
state_dict = {}
|
||||
|
||||
def update_sd(prefix, sd):
|
||||
for k, v in sd.items():
|
||||
key = prefix + k
|
||||
v = v.detach().clone().to("cpu").to(get_torch_dtype(save_dtype))
|
||||
state_dict[key] = v
|
||||
|
||||
# Convert the UNet model
|
||||
update_sd("model.diffusion_model.", self.unet.state_dict())
|
||||
|
||||
@@ -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)
|
||||
update_sd("conditioner.embedders.1.model.", text_enc2_dict)
|
||||
|
||||
else:
|
||||
# Convert the UNet model
|
||||
unet_state_dict = convert_unet_state_dict_to_sd(self.is_v2, self.unet.state_dict())
|
||||
update_sd("model.diffusion_model.", unet_state_dict)
|
||||
|
||||
# Convert the text encoder model
|
||||
if self.is_v2:
|
||||
make_dummy = True
|
||||
text_enc_dict = convert_text_encoder_state_dict_to_sd_v2(self.text_encoder.state_dict(), make_dummy)
|
||||
update_sd("cond_stage_model.model.", text_enc_dict)
|
||||
else:
|
||||
text_enc_dict = self.text_encoder.state_dict()
|
||||
update_sd("cond_stage_model.transformer.", text_enc_dict)
|
||||
|
||||
# Convert the VAE
|
||||
if self.vae is not None:
|
||||
vae_dict = model_util.convert_vae_state_dict(self.vae.state_dict())
|
||||
update_sd("first_stage_model.", vae_dict)
|
||||
|
||||
# Put together new checkpoint
|
||||
key_count = len(state_dict.keys())
|
||||
new_ckpt = {"state_dict": state_dict}
|
||||
|
||||
if model_util.is_safetensors(output_file):
|
||||
save_file(state_dict, output_file)
|
||||
else:
|
||||
torch.save(new_ckpt, output_file, meta)
|
||||
|
||||
return key_count
|
||||
else:
|
||||
raise NotImplementedError("sdv1.x, sdv2.x is not implemented yet")
|
||||
# prepare metadata
|
||||
meta = get_meta_for_safetensors(meta)
|
||||
# make sure parent folder exists
|
||||
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
||||
save_file(state_dict, output_file, metadata=meta)
|
||||
|
||||
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 hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
import sys
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from diffusers import (
|
||||
StableDiffusionPipeline,
|
||||
@@ -29,13 +34,16 @@ SCHEDLER_SCHEDULE = "scaled_linear"
|
||||
|
||||
|
||||
def get_torch_dtype(dtype_str):
|
||||
# if it is a torch dtype, return it
|
||||
if isinstance(dtype_str, torch.dtype):
|
||||
return dtype_str
|
||||
if dtype_str == "float" or dtype_str == "fp32" or dtype_str == "single" or dtype_str == "float32":
|
||||
return torch.float
|
||||
if dtype_str == "fp16" or dtype_str == "half" or dtype_str == "float16":
|
||||
return torch.float16
|
||||
if dtype_str == "bf16" or dtype_str == "bfloat16":
|
||||
return torch.bfloat16
|
||||
return None
|
||||
return dtype_str
|
||||
|
||||
|
||||
def replace_filewords_prompt(prompt, args: argparse.Namespace):
|
||||
@@ -399,3 +407,29 @@ def concat_prompt_embeddings(
|
||||
[unconditional.pooled_embeds, conditional.pooled_embeds]
|
||||
).repeat_interleave(n_imgs, dim=0)
|
||||
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