Compare commits
185 Commits
sdxl
...
developmen
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2a4b8e058 | ||
|
|
792a5e37e2 | ||
|
|
d7e55b6ad4 | ||
|
|
fbec68681d | ||
|
|
6280284d8b | ||
|
|
ad50921c41 | ||
|
|
e47006ed70 | ||
|
|
4f9cdd916a | ||
|
|
7782caa468 | ||
|
|
fa6d91ba76 | ||
|
|
1ee62562a4 | ||
|
|
dc8448d958 | ||
|
|
a8b3b8b8da | ||
|
|
93ea955d7c | ||
|
|
8a9e8f708f | ||
|
|
d35733ac06 | ||
|
|
ceaf1d9454 | ||
|
|
7d707b2fe6 | ||
|
|
a899ec91c8 | ||
|
|
436a09430e | ||
|
|
3097865203 | ||
|
|
b84e3260cb | ||
|
|
48a9bac22d | ||
|
|
298001439a | ||
|
|
6f3e0d5af2 | ||
|
|
0a79ac9604 | ||
|
|
9636194c09 | ||
|
|
d742792ee4 | ||
|
|
002279cec3 | ||
|
|
73c8b50975 | ||
|
|
34eb563d55 | ||
|
|
dc36bbb3c8 | ||
|
|
9905a1e205 | ||
|
|
0e9fc42816 | ||
|
|
d46112a354 | ||
|
|
07bf7bd7de | ||
|
|
da6302ada8 | ||
|
|
a05459afaf | ||
|
|
b1a22d0b3e | ||
|
|
7909b50d24 | ||
|
|
38e441a29c | ||
|
|
4e3b2c2569 | ||
|
|
239addba51 | ||
|
|
63ceffae24 | ||
|
|
f4c90bb589 | ||
|
|
bb1d3793e3 | ||
|
|
1d3de678aa | ||
|
|
b1cfafa0c6 | ||
|
|
cac8754399 | ||
|
|
f73402473b | ||
|
|
579650eaf8 | ||
|
|
320e109c5f | ||
|
|
560251a24f | ||
|
|
085787b799 | ||
|
|
8d9450ad7c | ||
|
|
8509da60cb | ||
|
|
c5d49ba661 | ||
|
|
76c764af49 | ||
|
|
abf7cd221d | ||
|
|
e5153d87c9 | ||
|
|
830e87cb87 | ||
|
|
19255cdc7c | ||
|
|
0f105690cc | ||
|
|
61badf85a7 | ||
|
|
181f237a7b | ||
|
|
c698837241 | ||
|
|
27f343fc08 | ||
|
|
17e4fe40d7 | ||
|
|
569d7464d5 | ||
|
|
4e945917df | ||
|
|
ae70200d3c | ||
|
|
d8d1e6fd1e | ||
|
|
257da9493d | ||
|
|
b5a2669b74 | ||
|
|
d74dd636ee | ||
|
|
e8583860ad | ||
|
|
083cefa78c | ||
|
|
b5ec8e4eb1 | ||
|
|
708b07adb7 | ||
|
|
a437aed45f | ||
|
|
34bfeba229 | ||
|
|
41a3f63b72 | ||
|
|
626ed2939a | ||
|
|
2128ac1e08 | ||
|
|
be804c9cf5 | ||
|
|
408c50ead1 | ||
|
|
4ed03a8d92 | ||
|
|
b01ab5d375 | ||
|
|
cb91b0d6da | ||
|
|
ce4f9fe02a | ||
|
|
92a086d5a5 | ||
|
|
3feb663a51 | ||
|
|
436bf0c6a3 | ||
|
|
f84500159c | ||
|
|
64a5441832 | ||
|
|
a4c3507a62 | ||
|
|
fa8fc32c0a | ||
|
|
22ed539321 | ||
|
|
7cd6945082 | ||
|
|
2a40937b4f | ||
|
|
4ca819a05e | ||
|
|
addf024630 | ||
|
|
33267e117c | ||
|
|
d401348c2e | ||
|
|
836fee47a6 | ||
|
|
14ff51ceb4 | ||
|
|
714854ee86 | ||
|
|
bd758ff203 | ||
|
|
a008d9e63b | ||
|
|
b79ced3e10 | ||
|
|
bee0b6a235 | ||
|
|
2ecb5cf024 | ||
|
|
fab7c2b04a | ||
|
|
e866c75638 | ||
|
|
71da78c8af | ||
|
|
c446f768ea | ||
|
|
cc49786ee9 | ||
|
|
9b164a8688 | ||
|
|
6bd3851058 | ||
|
|
fd338e67bb | ||
|
|
8105c05c12 | ||
|
|
2cb27c3f57 | ||
|
|
3367ab6b2c | ||
|
|
24f46ea7d6 | ||
|
|
5bef2985b5 | ||
|
|
aeaca13d69 | ||
|
|
b408f9f3eb | ||
|
|
7157c316af | ||
|
|
e2c547f6c2 | ||
|
|
7b770bc305 | ||
|
|
f200cf36c5 | ||
|
|
d298240cec | ||
|
|
2e6c55c720 | ||
|
|
36ba08d3fa | ||
|
|
e8667f856f | ||
|
|
bef5551ea5 | ||
|
|
b77b9acc0b | ||
|
|
c6675e2801 | ||
|
|
90eedb78bf | ||
|
|
80e2f4a2a4 | ||
|
|
d51c4ca704 | ||
|
|
c7ec132d5d | ||
|
|
ed9607e8da | ||
|
|
d44f8ac508 | ||
|
|
8d09eb44ec | ||
|
|
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
|
||||
6
.gitmodules
vendored
6
.gitmodules
vendored
@@ -4,3 +4,9 @@
|
||||
[submodule "repositories/leco"]
|
||||
path = repositories/leco
|
||||
url = https://github.com/p1atdev/LECO
|
||||
[submodule "repositories/batch_annotator"]
|
||||
path = repositories/batch_annotator
|
||||
url = https://github.com/ostris/batch-annotator
|
||||
[submodule "repositories/ipadapter"]
|
||||
path = repositories/ipadapter
|
||||
url = https://github.com/tencent-ailab/IP-Adapter.git
|
||||
|
||||
117
README.md
117
README.md
@@ -29,10 +29,24 @@ 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
|
||||
```
|
||||
|
||||
Windows:
|
||||
```bash
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
.\venv\Scripts\activate
|
||||
pip install torch --use-pep517 --extra-index-url https://download.pytorch.org/whl/cu118
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
|
||||
---
|
||||
|
||||
## Current Tools
|
||||
@@ -40,6 +54,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 +90,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 +144,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 +180,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"
|
||||
@@ -23,9 +23,8 @@ config:
|
||||
# network type lierla is traditional LoRA that works everywhere, only linear layers
|
||||
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
|
||||
|
||||
linear: 8
|
||||
linear_alpha: 4 # Do about half of rank
|
||||
# training config
|
||||
train:
|
||||
# this is also used in sampling. Stick with ddpm unless you know what you are doing
|
||||
@@ -33,14 +32,17 @@ 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
|
||||
# for sliders we are adjusting representation of the concept (unet),
|
||||
# not the description of it (text encoder)
|
||||
train_text_encoder: false
|
||||
|
||||
# same as from sd-scripts, not fully tested but should speed up training
|
||||
min_snr_gamma: 5.0
|
||||
# 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,14 +53,17 @@ 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
|
||||
# noise_offset: 0.0357 # SDXL was trained with offset of 0.0357. So use that when training on SDXL
|
||||
|
||||
# the model to train the LoRA network on
|
||||
model:
|
||||
@@ -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,48 @@ 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
|
||||
# shuffle the prompts split by the comma. We will run every combination randomly
|
||||
# this will make the LoRA more robust. You probably want this on unless prompt order
|
||||
# is important for some reason
|
||||
shuffle: true
|
||||
|
||||
# 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.
|
||||
102
extensions_built_in/advanced_generator/PureLoraGenerator.py
Normal file
102
extensions_built_in/advanced_generator/PureLoraGenerator.py
Normal file
@@ -0,0 +1,102 @@
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, SampleConfig, LoRMConfig
|
||||
from toolkit.lorm import ExtractMode, convert_diffusers_unet_to_lorm
|
||||
from toolkit.sd_device_states_presets import get_train_sd_device_state_preset
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class PureLoraGenerator(BaseExtensionProcess):
|
||||
|
||||
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.device = self.get_conf('device', 'cuda')
|
||||
self.device_torch = torch.device(self.device)
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.generate_config = SampleConfig(**self.get_conf('sample', required=True))
|
||||
self.dtype = self.get_conf('dtype', 'float16')
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
lorm_config = self.get_conf('lorm', None)
|
||||
self.lorm_config = LoRMConfig(**lorm_config) if lorm_config is not None else None
|
||||
|
||||
self.device_state_preset = get_train_sd_device_state_preset(
|
||||
device=torch.device(self.device),
|
||||
)
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
with torch.no_grad():
|
||||
self.sd.load_model()
|
||||
self.sd.unet.eval()
|
||||
self.sd.unet.to(self.device_torch)
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for te in self.sd.text_encoder:
|
||||
te.eval()
|
||||
te.to(self.device_torch)
|
||||
else:
|
||||
self.sd.text_encoder.eval()
|
||||
self.sd.to(self.device_torch)
|
||||
|
||||
print(f"Converting to LoRM UNet")
|
||||
# replace the unet with LoRMUnet
|
||||
convert_diffusers_unet_to_lorm(
|
||||
self.sd.unet,
|
||||
config=self.lorm_config,
|
||||
)
|
||||
|
||||
sample_folder = os.path.join(self.output_folder)
|
||||
gen_img_config_list = []
|
||||
|
||||
sample_config = self.generate_config
|
||||
start_seed = sample_config.seed
|
||||
current_seed = start_seed
|
||||
for i in range(len(sample_config.prompts)):
|
||||
if sample_config.walk_seed:
|
||||
current_seed = start_seed + i
|
||||
|
||||
filename = f"[time]_[count].{self.generate_config.ext}"
|
||||
output_path = os.path.join(sample_folder, filename)
|
||||
prompt = sample_config.prompts[i]
|
||||
extra_args = {}
|
||||
gen_img_config_list.append(GenerateImageConfig(
|
||||
prompt=prompt, # 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,
|
||||
output_ext=sample_config.ext,
|
||||
adapter_conditioning_scale=sample_config.adapter_conditioning_scale,
|
||||
**extra_args
|
||||
))
|
||||
|
||||
# send to be generated
|
||||
self.sd.generate_images(gen_img_config_list, sampler=sample_config.sampler)
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
193
extensions_built_in/advanced_generator/ReferenceGenerator.py
Normal file
193
extensions_built_in/advanced_generator/ReferenceGenerator.py
Normal file
@@ -0,0 +1,193 @@
|
||||
import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import List
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLAdapterPipeline
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
from toolkit.sampler import get_sampler
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.data_loader import get_dataloader_from_datasets
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from controlnet_aux.midas import MidasDetector
|
||||
from diffusers.utils import load_image
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.prompts: List[str]
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.walk_seed = kwargs.get('walk_seed', False)
|
||||
self.t2i_adapter_path = kwargs.get('t2i_adapter_path', None)
|
||||
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.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
|
||||
if kwargs.get('shuffle', False):
|
||||
# shuffle the prompts
|
||||
random.shuffle(self.prompts)
|
||||
|
||||
|
||||
class ReferenceGenerator(BaseExtensionProcess):
|
||||
|
||||
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.device = self.get_conf('device', 'cuda')
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
self.is_latents_cached = True
|
||||
raw_datasets = self.get_conf('datasets', None)
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
raw_datasets = preprocess_dataset_raw_config(raw_datasets)
|
||||
self.datasets = None
|
||||
self.datasets_reg = None
|
||||
self.dtype = self.get_conf('dtype', 'float16')
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.params = []
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
for raw_dataset in raw_datasets:
|
||||
dataset = DatasetConfig(**raw_dataset)
|
||||
is_caching = dataset.cache_latents or dataset.cache_latents_to_disk
|
||||
if not is_caching:
|
||||
self.is_latents_cached = False
|
||||
if dataset.is_reg:
|
||||
if self.datasets_reg is None:
|
||||
self.datasets_reg = []
|
||||
self.datasets_reg.append(dataset)
|
||||
else:
|
||||
if self.datasets is None:
|
||||
self.datasets = []
|
||||
self.datasets.append(dataset)
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
print(f"Using device {self.device}")
|
||||
self.data_loader: DataLoader = None
|
||||
self.adapter: T2IAdapter = None
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
device = torch.device(self.device)
|
||||
|
||||
if self.generate_config.t2i_adapter_path is not None:
|
||||
self.adapter = T2IAdapter.from_pretrained(
|
||||
"TencentARC/t2i-adapter-depth-midas-sdxl-1.0", torch_dtype=self.torch_dtype, varient="fp16"
|
||||
).to(device)
|
||||
|
||||
midas_depth = MidasDetector.from_pretrained(
|
||||
"valhalla/t2iadapter-aux-models", filename="dpt_large_384.pt", model_type="dpt_large"
|
||||
).to(device)
|
||||
|
||||
pipe = StableDiffusionXLAdapterPipeline(
|
||||
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=get_sampler(self.generate_config.sampler),
|
||||
adapter=self.adapter,
|
||||
).to(device)
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
|
||||
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
|
||||
|
||||
num_batches = len(self.data_loader)
|
||||
pbar = tqdm(total=num_batches, desc="Generating images")
|
||||
seed = self.generate_config.seed
|
||||
# load images from datasets, use tqdm
|
||||
for i, batch in enumerate(self.data_loader):
|
||||
batch: DataLoaderBatchDTO = batch
|
||||
|
||||
file_item: FileItemDTO = batch.file_items[0]
|
||||
img_path = file_item.path
|
||||
img_filename = os.path.basename(img_path)
|
||||
img_filename_no_ext = os.path.splitext(img_filename)[0]
|
||||
output_path = os.path.join(self.output_folder, img_filename)
|
||||
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
|
||||
output_depth_path = os.path.join(self.output_folder, img_filename_no_ext + '.depth.png')
|
||||
|
||||
caption = batch.get_caption_list()[0]
|
||||
|
||||
img: torch.Tensor = batch.tensor.clone()
|
||||
# image comes in -1 to 1. convert to a PIL RGB image
|
||||
img = (img + 1) / 2
|
||||
img = img.clamp(0, 1)
|
||||
img = img[0].permute(1, 2, 0).cpu().numpy()
|
||||
img = (img * 255).astype(np.uint8)
|
||||
image = Image.fromarray(img)
|
||||
|
||||
width, height = image.size
|
||||
min_res = min(width, height)
|
||||
|
||||
if self.generate_config.walk_seed:
|
||||
seed = seed + 1
|
||||
|
||||
if self.generate_config.seed == -1:
|
||||
# random
|
||||
seed = random.randint(0, 1000000)
|
||||
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
# generate depth map
|
||||
image = midas_depth(
|
||||
image,
|
||||
detect_resolution=min_res, # do 512 ?
|
||||
image_resolution=min_res
|
||||
)
|
||||
|
||||
# image.save(output_depth_path)
|
||||
|
||||
gen_images = pipe(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
image=image,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
adapter_conditioning_scale=self.generate_config.adapter_conditioning_scale,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
).images[0]
|
||||
gen_images.save(output_path)
|
||||
|
||||
# save caption
|
||||
with open(output_caption_path, 'w') as f:
|
||||
f.write(caption)
|
||||
|
||||
pbar.update(1)
|
||||
batch.cleanup()
|
||||
|
||||
pbar.close()
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
42
extensions_built_in/advanced_generator/__init__.py
Normal file
42
extensions_built_in/advanced_generator/__init__.py
Normal file
@@ -0,0 +1,42 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class AdvancedReferenceGeneratorExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "reference_generator"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Reference Generator"
|
||||
|
||||
# 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 .ReferenceGenerator import ReferenceGenerator
|
||||
return ReferenceGenerator
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class PureLoraGenerator(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "pure_lora_generator"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Pure LoRA Generator"
|
||||
|
||||
# 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 .PureLoraGenerator import PureLoraGenerator
|
||||
return PureLoraGenerator
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
AdvancedReferenceGeneratorExtension, PureLoraGenerator
|
||||
]
|
||||
@@ -0,0 +1,91 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: test_v1
|
||||
process:
|
||||
- type: 'textual_inversion_trainer'
|
||||
training_folder: "out/TI"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "out/.tensorboard"
|
||||
embedding:
|
||||
trigger: "your_trigger_here"
|
||||
tokens: 12
|
||||
init_words: "man with short brown hair"
|
||||
save_format: "safetensors" # 'safetensors' or 'pt'
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 100 # save every this many steps
|
||||
max_step_saves_to_keep: 5 # only affects step counts
|
||||
datasets:
|
||||
- folder_path: "/path/to/dataset"
|
||||
caption_ext: "txt"
|
||||
default_caption: "[trigger]"
|
||||
buckets: true
|
||||
resolution: 512
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 3000
|
||||
weight_jitter: 0.0
|
||||
lr: 5e-5
|
||||
train_unet: false
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: false
|
||||
optimizer: "adamw"
|
||||
# optimizer: "prodigy"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 4
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
min_snr_gamma: 5.0
|
||||
# skip_first_sample: true
|
||||
noise_offset: 0.0 # not needed for this
|
||||
model:
|
||||
# objective reality v2
|
||||
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of [trigger] laughing"
|
||||
- "photo of [trigger] smiling"
|
||||
- "[trigger] close up"
|
||||
- "dark scene [trigger] frozen"
|
||||
- "[trigger] nighttime"
|
||||
- "a painting of [trigger]"
|
||||
- "a drawing of [trigger]"
|
||||
- "a cartoon of [trigger]"
|
||||
- "[trigger] pixar style"
|
||||
- "[trigger] costume"
|
||||
neg: ""
|
||||
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
|
||||
|
||||
# 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.
|
||||
# It is saved in the model so be aware of that. The software will include this
|
||||
# plus some other information for you automatically
|
||||
meta:
|
||||
# [name] gets replaced with the name above
|
||||
name: "[name]"
|
||||
# version: '1.0'
|
||||
# creator:
|
||||
# name: Your Name
|
||||
# email: your@gmail.com
|
||||
# website: https://your.website
|
||||
151
extensions_built_in/concept_replacer/ConceptReplacer.py
Normal file
151
extensions_built_in/concept_replacer/ConceptReplacer.py
Normal file
@@ -0,0 +1,151 @@
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from torch.utils.data import DataLoader
|
||||
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, BlankNetwork
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class ConceptReplacementConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.concept: str = kwargs.get('concept', '')
|
||||
self.replacement: str = kwargs.get('replacement', '')
|
||||
|
||||
|
||||
class ConceptReplacer(BaseSDTrainProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
replacement_list = self.config.get('replacements', [])
|
||||
self.replacement_list = [ConceptReplacementConfig(**x) for x in replacement_list]
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
|
||||
# textual inversion
|
||||
if self.embedding is not None:
|
||||
# set text encoder to train. Not sure if this is necessary but diffusers example did it
|
||||
self.sd.text_encoder.train()
|
||||
|
||||
def hook_train_loop(self, batch):
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
|
||||
network_weight_list = batch.get_network_weight_list()
|
||||
|
||||
# have a blank network so we can wrap it in a context and set multipliers without checking every time
|
||||
if self.network is not None:
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
batch_replacement_list = []
|
||||
# get a random replacement for each prompt
|
||||
for prompt in conditioned_prompts:
|
||||
replacement = random.choice(self.replacement_list)
|
||||
batch_replacement_list.append(replacement)
|
||||
|
||||
# build out prompts
|
||||
concept_prompts = []
|
||||
replacement_prompts = []
|
||||
for idx, replacement in enumerate(batch_replacement_list):
|
||||
prompt = conditioned_prompts[idx]
|
||||
|
||||
# insert shuffled concept at beginning and end of prompt
|
||||
shuffled_concept = [x.strip() for x in replacement.concept.split(',')]
|
||||
random.shuffle(shuffled_concept)
|
||||
shuffled_concept = ', '.join(shuffled_concept)
|
||||
concept_prompts.append(f"{shuffled_concept}, {prompt}, {shuffled_concept}")
|
||||
|
||||
# insert replacement at beginning and end of prompt
|
||||
shuffled_replacement = [x.strip() for x in replacement.replacement.split(',')]
|
||||
random.shuffle(shuffled_replacement)
|
||||
shuffled_replacement = ', '.join(shuffled_replacement)
|
||||
replacement_prompts.append(f"{shuffled_replacement}, {prompt}, {shuffled_replacement}")
|
||||
|
||||
# predict the replacement without network
|
||||
conditional_embeds = self.sd.encode_prompt(replacement_prompts).to(self.device_torch, dtype=dtype)
|
||||
|
||||
replacement_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
del conditional_embeds
|
||||
replacement_pred = replacement_pred.detach()
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
flush()
|
||||
|
||||
# text encoding
|
||||
grad_on_text_encoder = False
|
||||
if self.train_config.train_text_encoder:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
if self.embedding:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
# set the weights
|
||||
network.multiplier = network_weight_list
|
||||
|
||||
# activate network if it exits
|
||||
with network:
|
||||
with torch.set_grad_enabled(grad_on_text_encoder):
|
||||
# embed the prompts
|
||||
conditional_embeds = self.sd.encode_prompt(concept_prompts).to(self.device_torch, dtype=dtype)
|
||||
if not grad_on_text_encoder:
|
||||
# detach the embeddings
|
||||
conditional_embeds = conditional_embeds.detach()
|
||||
self.optimizer.zero_grad()
|
||||
flush()
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
)
|
||||
|
||||
loss = torch.nn.functional.mse_loss(noise_pred.float(), replacement_pred.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler, self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean()
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
flush()
|
||||
|
||||
# apply gradients
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.embedding is not None:
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
self.embedding.restore_embeddings()
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss.item()}
|
||||
)
|
||||
# reset network multiplier
|
||||
network.multiplier = 1.0
|
||||
|
||||
return loss_dict
|
||||
26
extensions_built_in/concept_replacer/__init__.py
Normal file
26
extensions_built_in/concept_replacer/__init__.py
Normal file
@@ -0,0 +1,26 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class ConceptReplacerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "concept_replacer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Concept Replacer"
|
||||
|
||||
# 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 .ConceptReplacer import ConceptReplacer
|
||||
return ConceptReplacer
|
||||
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
ConceptReplacerExtension,
|
||||
]
|
||||
@@ -0,0 +1,91 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: test_v1
|
||||
process:
|
||||
- type: 'textual_inversion_trainer'
|
||||
training_folder: "out/TI"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "out/.tensorboard"
|
||||
embedding:
|
||||
trigger: "your_trigger_here"
|
||||
tokens: 12
|
||||
init_words: "man with short brown hair"
|
||||
save_format: "safetensors" # 'safetensors' or 'pt'
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 100 # save every this many steps
|
||||
max_step_saves_to_keep: 5 # only affects step counts
|
||||
datasets:
|
||||
- folder_path: "/path/to/dataset"
|
||||
caption_ext: "txt"
|
||||
default_caption: "[trigger]"
|
||||
buckets: true
|
||||
resolution: 512
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 3000
|
||||
weight_jitter: 0.0
|
||||
lr: 5e-5
|
||||
train_unet: false
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: false
|
||||
optimizer: "adamw"
|
||||
# optimizer: "prodigy"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 4
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
min_snr_gamma: 5.0
|
||||
# skip_first_sample: true
|
||||
noise_offset: 0.0 # not needed for this
|
||||
model:
|
||||
# objective reality v2
|
||||
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of [trigger] laughing"
|
||||
- "photo of [trigger] smiling"
|
||||
- "[trigger] close up"
|
||||
- "dark scene [trigger] frozen"
|
||||
- "[trigger] nighttime"
|
||||
- "a painting of [trigger]"
|
||||
- "a drawing of [trigger]"
|
||||
- "a cartoon of [trigger]"
|
||||
- "[trigger] pixar style"
|
||||
- "[trigger] costume"
|
||||
neg: ""
|
||||
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
|
||||
|
||||
# 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.
|
||||
# It is saved in the model so be aware of that. The software will include this
|
||||
# plus some other information for you automatically
|
||||
meta:
|
||||
# [name] gets replaced with the name above
|
||||
name: "[name]"
|
||||
# version: '1.0'
|
||||
# creator:
|
||||
# name: Your Name
|
||||
# email: your@gmail.com
|
||||
# website: https://your.website
|
||||
20
extensions_built_in/dataset_tools/DatasetTools.py
Normal file
20
extensions_built_in/dataset_tools/DatasetTools.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class DatasetTools(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
|
||||
raise NotImplementedError("This extension is not yet implemented")
|
||||
196
extensions_built_in/dataset_tools/SuperTagger.py
Normal file
196
extensions_built_in/dataset_tools/SuperTagger.py
Normal file
@@ -0,0 +1,196 @@
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
import traceback
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
from tqdm import tqdm
|
||||
|
||||
from .tools.dataset_tools_config_modules import RAW_DIR, TRAIN_DIR, Step, ImgInfo
|
||||
from .tools.fuyu_utils import FuyuImageProcessor
|
||||
from .tools.image_tools import load_image, ImageProcessor, resize_to_max
|
||||
from .tools.llava_utils import LLaVAImageProcessor
|
||||
from .tools.caption import default_long_prompt, default_short_prompt, default_replacements
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from .tools.sync_tools import get_img_paths
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
VERSION = 2
|
||||
|
||||
|
||||
class SuperTagger(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
parent_dir = config.get('parent_dir', None)
|
||||
self.dataset_paths: list[str] = config.get('dataset_paths', [])
|
||||
self.device = config.get('device', 'cuda')
|
||||
self.steps: list[Step] = config.get('steps', [])
|
||||
self.caption_method = config.get('caption_method', 'llava:default')
|
||||
self.caption_prompt = config.get('caption_prompt', default_long_prompt)
|
||||
self.caption_short_prompt = config.get('caption_short_prompt', default_short_prompt)
|
||||
self.force_reprocess_img = config.get('force_reprocess_img', False)
|
||||
self.caption_replacements = config.get('caption_replacements', default_replacements)
|
||||
self.caption_short_replacements = config.get('caption_short_replacements', default_replacements)
|
||||
self.master_dataset_dict = OrderedDict()
|
||||
self.dataset_master_config_file = config.get('dataset_master_config_file', None)
|
||||
if parent_dir is not None and len(self.dataset_paths) == 0:
|
||||
# find all folders in the patent_dataset_path
|
||||
self.dataset_paths = [
|
||||
os.path.join(parent_dir, folder)
|
||||
for folder in os.listdir(parent_dir)
|
||||
if os.path.isdir(os.path.join(parent_dir, folder))
|
||||
]
|
||||
else:
|
||||
# make sure they exist
|
||||
for dataset_path in self.dataset_paths:
|
||||
if not os.path.exists(dataset_path):
|
||||
raise ValueError(f"Dataset path does not exist: {dataset_path}")
|
||||
|
||||
print(f"Found {len(self.dataset_paths)} dataset paths")
|
||||
|
||||
self.image_processor: ImageProcessor = self.get_image_processor()
|
||||
|
||||
def get_image_processor(self):
|
||||
if self.caption_method.startswith('llava'):
|
||||
return LLaVAImageProcessor(device=self.device)
|
||||
elif self.caption_method.startswith('fuyu'):
|
||||
return FuyuImageProcessor(device=self.device)
|
||||
else:
|
||||
raise ValueError(f"Unknown caption method: {self.caption_method}")
|
||||
|
||||
def process_image(self, img_path: str):
|
||||
root_img_dir = os.path.dirname(os.path.dirname(img_path))
|
||||
filename = os.path.basename(img_path)
|
||||
filename_no_ext = os.path.splitext(filename)[0]
|
||||
train_dir = os.path.join(root_img_dir, TRAIN_DIR)
|
||||
train_img_path = os.path.join(train_dir, filename)
|
||||
json_path = os.path.join(train_dir, f"{filename_no_ext}.json")
|
||||
|
||||
# check if json exists, if it does load it as image info
|
||||
if os.path.exists(json_path):
|
||||
with open(json_path, 'r') as f:
|
||||
img_info = ImgInfo(**json.load(f))
|
||||
else:
|
||||
img_info = ImgInfo()
|
||||
|
||||
# always send steps first in case other processes need them
|
||||
img_info.add_steps(copy.deepcopy(self.steps))
|
||||
img_info.set_version(VERSION)
|
||||
img_info.set_caption_method(self.caption_method)
|
||||
|
||||
image: Image = None
|
||||
caption_image: Image = None
|
||||
|
||||
did_update_image = False
|
||||
|
||||
# trigger reprocess of steps
|
||||
if self.force_reprocess_img:
|
||||
img_info.trigger_image_reprocess()
|
||||
|
||||
# set the image as updated if it does not exist on disk
|
||||
if not os.path.exists(train_img_path):
|
||||
did_update_image = True
|
||||
image = load_image(img_path)
|
||||
if img_info.force_image_process:
|
||||
did_update_image = True
|
||||
image = load_image(img_path)
|
||||
|
||||
# go through the needed steps
|
||||
for step in copy.deepcopy(img_info.state.steps_to_complete):
|
||||
if step == 'caption':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
if caption_image is None:
|
||||
caption_image = resize_to_max(image, 1024, 1024)
|
||||
|
||||
if not self.image_processor.is_loaded:
|
||||
print('Loading Model. Takes a while, especially the first time')
|
||||
self.image_processor.load_model()
|
||||
|
||||
img_info.caption = self.image_processor.generate_caption(
|
||||
image=caption_image,
|
||||
prompt=self.caption_prompt,
|
||||
replacements=self.caption_replacements
|
||||
)
|
||||
img_info.mark_step_complete(step)
|
||||
elif step == 'caption_short':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
|
||||
if caption_image is None:
|
||||
caption_image = resize_to_max(image, 1024, 1024)
|
||||
|
||||
if not self.image_processor.is_loaded:
|
||||
print('Loading Model. Takes a while, especially the first time')
|
||||
self.image_processor.load_model()
|
||||
img_info.caption_short = self.image_processor.generate_caption(
|
||||
image=caption_image,
|
||||
prompt=self.caption_short_prompt,
|
||||
replacements=self.caption_short_replacements
|
||||
)
|
||||
img_info.mark_step_complete(step)
|
||||
elif step == 'contrast_stretch':
|
||||
# load image
|
||||
if image is None:
|
||||
image = load_image(img_path)
|
||||
image = ImageOps.autocontrast(image, cutoff=(0.1, 0), preserve_tone=True)
|
||||
did_update_image = True
|
||||
img_info.mark_step_complete(step)
|
||||
else:
|
||||
raise ValueError(f"Unknown step: {step}")
|
||||
|
||||
os.makedirs(os.path.dirname(train_img_path), exist_ok=True)
|
||||
if did_update_image:
|
||||
image.save(train_img_path)
|
||||
|
||||
if img_info.is_dirty:
|
||||
with open(json_path, 'w') as f:
|
||||
json.dump(img_info.to_dict(), f, indent=4)
|
||||
|
||||
if self.dataset_master_config_file:
|
||||
# add to master dict
|
||||
self.master_dataset_dict[train_img_path] = img_info.to_dict()
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
imgs_to_process = []
|
||||
# find all images
|
||||
for dataset_path in self.dataset_paths:
|
||||
raw_dir = os.path.join(dataset_path, RAW_DIR)
|
||||
raw_image_paths = get_img_paths(raw_dir)
|
||||
for raw_image_path in raw_image_paths:
|
||||
imgs_to_process.append(raw_image_path)
|
||||
|
||||
if len(imgs_to_process) == 0:
|
||||
print(f"No images to process")
|
||||
else:
|
||||
print(f"Found {len(imgs_to_process)} to process")
|
||||
|
||||
for img_path in tqdm(imgs_to_process, desc="Processing images"):
|
||||
try:
|
||||
self.process_image(img_path)
|
||||
except Exception:
|
||||
# print full stack trace
|
||||
print(traceback.format_exc())
|
||||
continue
|
||||
# self.process_image(img_path)
|
||||
|
||||
if self.dataset_master_config_file is not None:
|
||||
# save it as json
|
||||
with open(self.dataset_master_config_file, 'w') as f:
|
||||
json.dump(self.master_dataset_dict, f, indent=4)
|
||||
|
||||
del self.image_processor
|
||||
flush()
|
||||
131
extensions_built_in/dataset_tools/SyncFromCollection.py
Normal file
131
extensions_built_in/dataset_tools/SyncFromCollection.py
Normal file
@@ -0,0 +1,131 @@
|
||||
import os
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
import gc
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from .tools.dataset_tools_config_modules import DatasetSyncCollectionConfig, RAW_DIR, NEW_DIR
|
||||
from .tools.sync_tools import get_unsplash_images, get_pexels_images, get_local_image_file_names, download_image, \
|
||||
get_img_paths
|
||||
from jobs.process import BaseExtensionProcess
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class SyncFromCollection(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
|
||||
self.min_width = config.get('min_width', 1024)
|
||||
self.min_height = config.get('min_height', 1024)
|
||||
|
||||
# add our min_width and min_height to each dataset config if they don't exist
|
||||
for dataset_config in config.get('dataset_sync', []):
|
||||
if 'min_width' not in dataset_config:
|
||||
dataset_config['min_width'] = self.min_width
|
||||
if 'min_height' not in dataset_config:
|
||||
dataset_config['min_height'] = self.min_height
|
||||
|
||||
self.dataset_configs: List[DatasetSyncCollectionConfig] = [
|
||||
DatasetSyncCollectionConfig(**dataset_config)
|
||||
for dataset_config in config.get('dataset_sync', [])
|
||||
]
|
||||
print(f"Found {len(self.dataset_configs)} dataset configs")
|
||||
|
||||
def move_new_images(self, root_dir: str):
|
||||
raw_dir = os.path.join(root_dir, RAW_DIR)
|
||||
new_dir = os.path.join(root_dir, NEW_DIR)
|
||||
new_images = get_img_paths(new_dir)
|
||||
|
||||
for img_path in new_images:
|
||||
# move to raw
|
||||
new_path = os.path.join(raw_dir, os.path.basename(img_path))
|
||||
shutil.move(img_path, new_path)
|
||||
|
||||
# remove new dir
|
||||
shutil.rmtree(new_dir)
|
||||
|
||||
def sync_dataset(self, config: DatasetSyncCollectionConfig):
|
||||
if config.host == 'unsplash':
|
||||
get_images = get_unsplash_images
|
||||
elif config.host == 'pexels':
|
||||
get_images = get_pexels_images
|
||||
else:
|
||||
raise ValueError(f"Unknown host: {config.host}")
|
||||
|
||||
results = {
|
||||
'num_downloaded': 0,
|
||||
'num_skipped': 0,
|
||||
'bad': 0,
|
||||
'total': 0,
|
||||
}
|
||||
|
||||
photos = get_images(config)
|
||||
raw_dir = os.path.join(config.directory, RAW_DIR)
|
||||
new_dir = os.path.join(config.directory, NEW_DIR)
|
||||
raw_images = get_local_image_file_names(raw_dir)
|
||||
new_images = get_local_image_file_names(new_dir)
|
||||
|
||||
for photo in tqdm(photos, desc=f"{config.host}-{config.collection_id}"):
|
||||
try:
|
||||
if photo.filename not in raw_images and photo.filename not in new_images:
|
||||
download_image(photo, new_dir, min_width=self.min_width, min_height=self.min_height)
|
||||
results['num_downloaded'] += 1
|
||||
else:
|
||||
results['num_skipped'] += 1
|
||||
except Exception as e:
|
||||
print(f" - BAD({photo.id}): {e}")
|
||||
results['bad'] += 1
|
||||
continue
|
||||
results['total'] += 1
|
||||
|
||||
return results
|
||||
|
||||
def print_results(self, results):
|
||||
print(
|
||||
f" - new:{results['num_downloaded']}, old:{results['num_skipped']}, bad:{results['bad']} total:{results['total']}")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print(f"Syncing {len(self.dataset_configs)} datasets")
|
||||
all_results = None
|
||||
failed_datasets = []
|
||||
for dataset_config in tqdm(self.dataset_configs, desc="Syncing datasets", leave=True):
|
||||
try:
|
||||
results = self.sync_dataset(dataset_config)
|
||||
if all_results is None:
|
||||
all_results = {**results}
|
||||
else:
|
||||
for key, value in results.items():
|
||||
all_results[key] += value
|
||||
|
||||
self.print_results(results)
|
||||
except Exception as e:
|
||||
print(f" - FAILED: {e}")
|
||||
if 'response' in e.__dict__:
|
||||
error = f"{e.response.status_code}: {e.response.text}"
|
||||
print(f" - {error}")
|
||||
failed_datasets.append({'dataset': dataset_config, 'error': error})
|
||||
else:
|
||||
failed_datasets.append({'dataset': dataset_config, 'error': str(e)})
|
||||
continue
|
||||
|
||||
print("Moving new images to raw")
|
||||
for dataset_config in self.dataset_configs:
|
||||
self.move_new_images(dataset_config.directory)
|
||||
|
||||
print("Done syncing datasets")
|
||||
self.print_results(all_results)
|
||||
|
||||
if len(failed_datasets) > 0:
|
||||
print(f"Failed to sync {len(failed_datasets)} datasets")
|
||||
for failed in failed_datasets:
|
||||
print(f" - {failed['dataset'].host}-{failed['dataset'].collection_id}")
|
||||
print(f" - ERR: {failed['error']}")
|
||||
43
extensions_built_in/dataset_tools/__init__.py
Normal file
43
extensions_built_in/dataset_tools/__init__.py
Normal file
@@ -0,0 +1,43 @@
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
class DatasetToolsExtension(Extension):
|
||||
uid = "dataset_tools"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Dataset Tools"
|
||||
|
||||
# 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 .DatasetTools import DatasetTools
|
||||
return DatasetTools
|
||||
|
||||
|
||||
class SyncFromCollectionExtension(Extension):
|
||||
uid = "sync_from_collection"
|
||||
name = "Sync from Collection"
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .SyncFromCollection import SyncFromCollection
|
||||
return SyncFromCollection
|
||||
|
||||
|
||||
class SuperTaggerExtension(Extension):
|
||||
uid = "super_tagger"
|
||||
name = "Super Tagger"
|
||||
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .SuperTagger import SuperTagger
|
||||
return SuperTagger
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
SyncFromCollectionExtension, DatasetToolsExtension, SuperTaggerExtension
|
||||
]
|
||||
53
extensions_built_in/dataset_tools/tools/caption.py
Normal file
53
extensions_built_in/dataset_tools/tools/caption.py
Normal file
@@ -0,0 +1,53 @@
|
||||
|
||||
caption_manipulation_steps = ['caption', 'caption_short']
|
||||
|
||||
default_long_prompt = 'caption this image. describe every single thing in the image in detail. Do not include any unnecessary words in your description for the sake of good grammar. I want many short statements that serve the single purpose of giving the most thorough description if items as possible in the smallest, comma separated way possible. be sure to describe people\'s moods, clothing, the environment, lighting, colors, and everything.'
|
||||
default_short_prompt = 'caption this image in less than ten words'
|
||||
|
||||
default_replacements = [
|
||||
("the image features", ""),
|
||||
("the image shows", ""),
|
||||
("the image depicts", ""),
|
||||
("the image is", ""),
|
||||
("in this image", ""),
|
||||
("in the image", ""),
|
||||
]
|
||||
|
||||
|
||||
def clean_caption(cap, replacements=None):
|
||||
if replacements is None:
|
||||
replacements = default_replacements
|
||||
|
||||
# remove any newlines
|
||||
cap = cap.replace("\n", ", ")
|
||||
cap = cap.replace("\r", ", ")
|
||||
cap = cap.replace(".", ",")
|
||||
cap = cap.replace("\"", "")
|
||||
|
||||
# remove unicode characters
|
||||
cap = cap.encode('ascii', 'ignore').decode('ascii')
|
||||
|
||||
# make lowercase
|
||||
cap = cap.lower()
|
||||
# remove any extra spaces
|
||||
cap = " ".join(cap.split())
|
||||
|
||||
for replacement in replacements:
|
||||
if replacement[0].startswith('*'):
|
||||
# we are removing all text if it starts with this and the rest matches
|
||||
search_text = replacement[0][1:]
|
||||
if cap.startswith(search_text):
|
||||
cap = ""
|
||||
else:
|
||||
cap = cap.replace(replacement[0].lower(), replacement[1].lower())
|
||||
|
||||
cap_list = cap.split(",")
|
||||
# trim whitespace
|
||||
cap_list = [c.strip() for c in cap_list]
|
||||
# remove empty strings
|
||||
cap_list = [c for c in cap_list if c != ""]
|
||||
# remove duplicates
|
||||
cap_list = list(dict.fromkeys(cap_list))
|
||||
# join back together
|
||||
cap = ", ".join(cap_list)
|
||||
return cap
|
||||
@@ -0,0 +1,187 @@
|
||||
import json
|
||||
from typing import Literal, Type, TYPE_CHECKING
|
||||
|
||||
Host: Type = Literal['unsplash', 'pexels']
|
||||
|
||||
RAW_DIR = "raw"
|
||||
NEW_DIR = "_tmp"
|
||||
TRAIN_DIR = "train"
|
||||
DEPTH_DIR = "depth"
|
||||
|
||||
from .image_tools import Step, img_manipulation_steps
|
||||
from .caption import caption_manipulation_steps
|
||||
|
||||
|
||||
class DatasetSyncCollectionConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.host: Host = kwargs.get('host', None)
|
||||
self.collection_id: str = kwargs.get('collection_id', None)
|
||||
self.directory: str = kwargs.get('directory', None)
|
||||
self.api_key: str = kwargs.get('api_key', None)
|
||||
self.min_width: int = kwargs.get('min_width', 1024)
|
||||
self.min_height: int = kwargs.get('min_height', 1024)
|
||||
|
||||
if self.host is None:
|
||||
raise ValueError("host is required")
|
||||
if self.collection_id is None:
|
||||
raise ValueError("collection_id is required")
|
||||
if self.directory is None:
|
||||
raise ValueError("directory is required")
|
||||
if self.api_key is None:
|
||||
raise ValueError(f"api_key is required: {self.host}:{self.collection_id}")
|
||||
|
||||
|
||||
class ImageState:
|
||||
def __init__(self, **kwargs):
|
||||
self.steps_complete: list[Step] = kwargs.get('steps_complete', [])
|
||||
self.steps_to_complete: list[Step] = kwargs.get('steps_to_complete', [])
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'steps_complete': self.steps_complete
|
||||
}
|
||||
|
||||
|
||||
class Rect:
|
||||
def __init__(self, **kwargs):
|
||||
self.x = kwargs.get('x', 0)
|
||||
self.y = kwargs.get('y', 0)
|
||||
self.width = kwargs.get('width', 0)
|
||||
self.height = kwargs.get('height', 0)
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'x': self.x,
|
||||
'y': self.y,
|
||||
'width': self.width,
|
||||
'height': self.height
|
||||
}
|
||||
|
||||
|
||||
class ImgInfo:
|
||||
def __init__(self, **kwargs):
|
||||
self.version: int = kwargs.get('version', None)
|
||||
self.caption: str = kwargs.get('caption', None)
|
||||
self.caption_short: str = kwargs.get('caption_short', None)
|
||||
self.poi = [Rect(**poi) for poi in kwargs.get('poi', [])]
|
||||
self.state = ImageState(**kwargs.get('state', {}))
|
||||
self.caption_method = kwargs.get('caption_method', None)
|
||||
self.other_captions = kwargs.get('other_captions', {})
|
||||
self._upgrade_state()
|
||||
self.force_image_process: bool = False
|
||||
self._requested_steps: list[Step] = []
|
||||
|
||||
self.is_dirty: bool = False
|
||||
|
||||
def _upgrade_state(self):
|
||||
# upgrades older states
|
||||
if self.caption is not None and 'caption' not in self.state.steps_complete:
|
||||
self.mark_step_complete('caption')
|
||||
self.is_dirty = True
|
||||
if self.caption_short is not None and 'caption_short' not in self.state.steps_complete:
|
||||
self.mark_step_complete('caption_short')
|
||||
self.is_dirty = True
|
||||
if self.caption_method is None and self.caption is not None:
|
||||
# added caption method in version 2. Was all llava before that
|
||||
self.caption_method = 'llava:default'
|
||||
self.is_dirty = True
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
'version': self.version,
|
||||
'caption_method': self.caption_method,
|
||||
'caption': self.caption,
|
||||
'caption_short': self.caption_short,
|
||||
'poi': [poi.to_dict() for poi in self.poi],
|
||||
'state': self.state.to_dict(),
|
||||
'other_captions': self.other_captions
|
||||
}
|
||||
|
||||
def mark_step_complete(self, step: Step):
|
||||
if step not in self.state.steps_complete:
|
||||
self.state.steps_complete.append(step)
|
||||
if step in self.state.steps_to_complete:
|
||||
self.state.steps_to_complete.remove(step)
|
||||
self.is_dirty = True
|
||||
|
||||
def add_step(self, step: Step):
|
||||
if step not in self.state.steps_to_complete and step not in self.state.steps_complete:
|
||||
self.state.steps_to_complete.append(step)
|
||||
|
||||
def trigger_image_reprocess(self):
|
||||
if self._requested_steps is None:
|
||||
raise Exception("Must call add_steps before trigger_image_reprocess")
|
||||
steps = self._requested_steps
|
||||
# remove all image manipulationf from steps_to_complete
|
||||
for step in img_manipulation_steps:
|
||||
if step in self.state.steps_to_complete:
|
||||
self.state.steps_to_complete.remove(step)
|
||||
if step in self.state.steps_complete:
|
||||
self.state.steps_complete.remove(step)
|
||||
self.force_image_process = True
|
||||
self.is_dirty = True
|
||||
# we want to keep the order passed in process file
|
||||
for step in steps:
|
||||
if step in img_manipulation_steps:
|
||||
self.add_step(step)
|
||||
|
||||
def add_steps(self, steps: list[Step]):
|
||||
self._requested_steps = [step for step in steps]
|
||||
for stage in steps:
|
||||
self.add_step(stage)
|
||||
|
||||
# update steps if we have any img processes not complete, we have to reprocess them all
|
||||
# if any steps_to_complete are in img_manipulation_steps
|
||||
|
||||
is_manipulating_image = any([step in img_manipulation_steps for step in self.state.steps_to_complete])
|
||||
order_has_changed = False
|
||||
|
||||
if not is_manipulating_image:
|
||||
# check to see if order has changed. No need to if already redoing it. Will detect if ones are removed
|
||||
target_img_manipulation_order = [step for step in steps if step in img_manipulation_steps]
|
||||
current_img_manipulation_order = [step for step in self.state.steps_complete if
|
||||
step in img_manipulation_steps]
|
||||
if target_img_manipulation_order != current_img_manipulation_order:
|
||||
order_has_changed = True
|
||||
|
||||
if is_manipulating_image or order_has_changed:
|
||||
self.trigger_image_reprocess()
|
||||
|
||||
def set_caption_method(self, method: str):
|
||||
if self._requested_steps is None:
|
||||
raise Exception("Must call add_steps before set_caption_method")
|
||||
if self.caption_method != method:
|
||||
self.is_dirty = True
|
||||
# move previous caption method to other_captions
|
||||
if self.caption_method is not None and self.caption is not None or self.caption_short is not None:
|
||||
self.other_captions[self.caption_method] = {
|
||||
'caption': self.caption,
|
||||
'caption_short': self.caption_short,
|
||||
}
|
||||
self.caption_method = method
|
||||
self.caption = None
|
||||
self.caption_short = None
|
||||
# see if we have a caption from the new method
|
||||
if method in self.other_captions:
|
||||
self.caption = self.other_captions[method].get('caption', None)
|
||||
self.caption_short = self.other_captions[method].get('caption_short', None)
|
||||
else:
|
||||
self.trigger_new_caption()
|
||||
|
||||
def trigger_new_caption(self):
|
||||
self.caption = None
|
||||
self.caption_short = None
|
||||
self.is_dirty = True
|
||||
# check to see if we have any steps in the complete list and move them to the to_complete list
|
||||
for step in self.state.steps_complete:
|
||||
if step in caption_manipulation_steps:
|
||||
self.state.steps_complete.remove(step)
|
||||
self.state.steps_to_complete.append(step)
|
||||
|
||||
def to_json(self):
|
||||
return json.dumps(self.to_dict())
|
||||
|
||||
def set_version(self, version: int):
|
||||
if self.version != version:
|
||||
self.is_dirty = True
|
||||
self.version = version
|
||||
66
extensions_built_in/dataset_tools/tools/fuyu_utils.py
Normal file
66
extensions_built_in/dataset_tools/tools/fuyu_utils.py
Normal file
@@ -0,0 +1,66 @@
|
||||
from transformers import CLIPImageProcessor, BitsAndBytesConfig, AutoTokenizer
|
||||
|
||||
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class FuyuImageProcessor:
|
||||
def __init__(self, device='cuda'):
|
||||
from transformers import FuyuProcessor, FuyuForCausalLM
|
||||
self.device = device
|
||||
self.model: FuyuForCausalLM = None
|
||||
self.processor: FuyuProcessor = None
|
||||
self.dtype = torch.bfloat16
|
||||
self.tokenizer: AutoTokenizer
|
||||
self.is_loaded = False
|
||||
|
||||
def load_model(self):
|
||||
from transformers import FuyuProcessor, FuyuForCausalLM
|
||||
model_path = "adept/fuyu-8b"
|
||||
kwargs = {"device_map": self.device}
|
||||
kwargs['load_in_4bit'] = True
|
||||
kwargs['quantization_config'] = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=self.dtype,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_quant_type='nf4'
|
||||
)
|
||||
self.processor = FuyuProcessor.from_pretrained(model_path)
|
||||
self.model = FuyuForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
|
||||
self.is_loaded = True
|
||||
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
self.model = FuyuForCausalLM.from_pretrained(model_path, torch_dtype=self.dtype, **kwargs)
|
||||
self.processor = FuyuProcessor(image_processor=FuyuImageProcessor(), tokenizer=self.tokenizer)
|
||||
|
||||
def generate_caption(
|
||||
self, image: Image,
|
||||
prompt: str = default_long_prompt,
|
||||
replacements=default_replacements,
|
||||
max_new_tokens=512
|
||||
):
|
||||
# prepare inputs for the model
|
||||
# text_prompt = f"{prompt}\n"
|
||||
|
||||
# image = image.convert('RGB')
|
||||
model_inputs = self.processor(text=prompt, images=[image])
|
||||
model_inputs = {k: v.to(dtype=self.dtype if torch.is_floating_point(v) else v.dtype, device=self.device) for k, v in
|
||||
model_inputs.items()}
|
||||
|
||||
generation_output = self.model.generate(**model_inputs, max_new_tokens=max_new_tokens)
|
||||
prompt_len = model_inputs["input_ids"].shape[-1]
|
||||
output = self.tokenizer.decode(generation_output[0][prompt_len:], skip_special_tokens=True)
|
||||
output = clean_caption(output, replacements=replacements)
|
||||
return output
|
||||
|
||||
# inputs = self.processor(text=text_prompt, images=image, return_tensors="pt")
|
||||
# for k, v in inputs.items():
|
||||
# inputs[k] = v.to(self.device)
|
||||
|
||||
# # autoregressively generate text
|
||||
# generation_output = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
|
||||
# generation_text = self.processor.batch_decode(generation_output[:, -max_new_tokens:], skip_special_tokens=True)
|
||||
# output = generation_text[0]
|
||||
#
|
||||
# return clean_caption(output, replacements=replacements)
|
||||
49
extensions_built_in/dataset_tools/tools/image_tools.py
Normal file
49
extensions_built_in/dataset_tools/tools/image_tools.py
Normal file
@@ -0,0 +1,49 @@
|
||||
from typing import Literal, Type, TYPE_CHECKING, Union
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
Step: Type = Literal['caption', 'caption_short', 'create_mask', 'contrast_stretch']
|
||||
|
||||
img_manipulation_steps = ['contrast_stretch']
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .llava_utils import LLaVAImageProcessor
|
||||
from .fuyu_utils import FuyuImageProcessor
|
||||
|
||||
ImageProcessor = Union['LLaVAImageProcessor', 'FuyuImageProcessor']
|
||||
|
||||
|
||||
def pil_to_cv2(image):
|
||||
"""Convert a PIL image to a cv2 image."""
|
||||
return cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
|
||||
|
||||
|
||||
def cv2_to_pil(image):
|
||||
"""Convert a cv2 image to a PIL image."""
|
||||
return Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
|
||||
|
||||
|
||||
def load_image(img_path: str):
|
||||
image = Image.open(img_path).convert('RGB')
|
||||
try:
|
||||
# transpose with exif data
|
||||
image = ImageOps.exif_transpose(image)
|
||||
except Exception as e:
|
||||
pass
|
||||
return image
|
||||
|
||||
|
||||
def resize_to_max(image, max_width=1024, max_height=1024):
|
||||
width, height = image.size
|
||||
if width <= max_width and height <= max_height:
|
||||
return image
|
||||
|
||||
scale = min(max_width / width, max_height / height)
|
||||
width = int(width * scale)
|
||||
height = int(height * scale)
|
||||
|
||||
return image.resize((width, height), Image.LANCZOS)
|
||||
85
extensions_built_in/dataset_tools/tools/llava_utils.py
Normal file
85
extensions_built_in/dataset_tools/tools/llava_utils.py
Normal file
@@ -0,0 +1,85 @@
|
||||
|
||||
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
|
||||
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
from transformers import AutoTokenizer, BitsAndBytesConfig, CLIPImageProcessor
|
||||
|
||||
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
|
||||
class LLaVAImageProcessor:
|
||||
def __init__(self, device='cuda'):
|
||||
try:
|
||||
from llava.model import LlavaLlamaForCausalLM
|
||||
except ImportError:
|
||||
# print("You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
|
||||
print(
|
||||
"You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
|
||||
raise
|
||||
self.device = device
|
||||
self.model: LlavaLlamaForCausalLM = None
|
||||
self.tokenizer: AutoTokenizer = None
|
||||
self.image_processor: CLIPImageProcessor = None
|
||||
self.is_loaded = False
|
||||
|
||||
def load_model(self):
|
||||
from llava.model import LlavaLlamaForCausalLM
|
||||
|
||||
model_path = "4bit/llava-v1.5-13b-3GB"
|
||||
# kwargs = {"device_map": "auto"}
|
||||
kwargs = {"device_map": self.device}
|
||||
kwargs['load_in_4bit'] = True
|
||||
kwargs['quantization_config'] = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.float16,
|
||||
bnb_4bit_use_double_quant=True,
|
||||
bnb_4bit_quant_type='nf4'
|
||||
)
|
||||
self.model = LlavaLlamaForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=False)
|
||||
vision_tower = self.model.get_vision_tower()
|
||||
if not vision_tower.is_loaded:
|
||||
vision_tower.load_model()
|
||||
vision_tower.to(device=self.device)
|
||||
self.image_processor = vision_tower.image_processor
|
||||
self.is_loaded = True
|
||||
|
||||
def generate_caption(
|
||||
self, image:
|
||||
Image, prompt: str = default_long_prompt,
|
||||
replacements=default_replacements,
|
||||
max_new_tokens=512
|
||||
):
|
||||
from llava.conversation import conv_templates, SeparatorStyle
|
||||
from llava.utils import disable_torch_init
|
||||
from llava.constants import IMAGE_TOKEN_INDEX, DEFAULT_IMAGE_TOKEN, DEFAULT_IM_START_TOKEN, DEFAULT_IM_END_TOKEN
|
||||
from llava.mm_utils import tokenizer_image_token, KeywordsStoppingCriteria
|
||||
# question = "how many dogs are in the picture?"
|
||||
disable_torch_init()
|
||||
conv_mode = "llava_v0"
|
||||
conv = conv_templates[conv_mode].copy()
|
||||
roles = conv.roles
|
||||
image_tensor = self.image_processor.preprocess([image], return_tensors='pt')['pixel_values'].half().cuda()
|
||||
|
||||
inp = f"{roles[0]}: {prompt}"
|
||||
inp = DEFAULT_IM_START_TOKEN + DEFAULT_IMAGE_TOKEN + DEFAULT_IM_END_TOKEN + '\n' + inp
|
||||
conv.append_message(conv.roles[0], inp)
|
||||
conv.append_message(conv.roles[1], None)
|
||||
raw_prompt = conv.get_prompt()
|
||||
input_ids = tokenizer_image_token(raw_prompt, self.tokenizer, IMAGE_TOKEN_INDEX,
|
||||
return_tensors='pt').unsqueeze(0).cuda()
|
||||
stop_str = conv.sep if conv.sep_style != SeparatorStyle.TWO else conv.sep2
|
||||
keywords = [stop_str]
|
||||
stopping_criteria = KeywordsStoppingCriteria(keywords, self.tokenizer, input_ids)
|
||||
with torch.inference_mode():
|
||||
output_ids = self.model.generate(
|
||||
input_ids, images=image_tensor, do_sample=True, temperature=0.1,
|
||||
max_new_tokens=max_new_tokens, use_cache=True, stopping_criteria=[stopping_criteria],
|
||||
top_p=0.8
|
||||
)
|
||||
outputs = self.tokenizer.decode(output_ids[0, input_ids.shape[1]:]).strip()
|
||||
conv.messages[-1][-1] = outputs
|
||||
output = outputs.rsplit('</s>', 1)[0]
|
||||
return clean_caption(output, replacements=replacements)
|
||||
279
extensions_built_in/dataset_tools/tools/sync_tools.py
Normal file
279
extensions_built_in/dataset_tools/tools/sync_tools.py
Normal file
@@ -0,0 +1,279 @@
|
||||
import os
|
||||
import requests
|
||||
import tqdm
|
||||
from typing import List, Optional, TYPE_CHECKING
|
||||
|
||||
|
||||
def img_root_path(img_id: str):
|
||||
return os.path.dirname(os.path.dirname(img_id))
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .dataset_tools_config_modules import DatasetSyncCollectionConfig
|
||||
|
||||
img_exts = ['.jpg', '.jpeg', '.webp', '.png']
|
||||
|
||||
class Photo:
|
||||
def __init__(
|
||||
self,
|
||||
id,
|
||||
host,
|
||||
width,
|
||||
height,
|
||||
url,
|
||||
filename
|
||||
):
|
||||
self.id = str(id)
|
||||
self.host = host
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.url = url
|
||||
self.filename = filename
|
||||
|
||||
|
||||
def get_desired_size(img_width: int, img_height: int, min_width: int, min_height: int):
|
||||
if img_width > img_height:
|
||||
scale = min_height / img_height
|
||||
else:
|
||||
scale = min_width / img_width
|
||||
|
||||
new_width = int(img_width * scale)
|
||||
new_height = int(img_height * scale)
|
||||
|
||||
return new_width, new_height
|
||||
|
||||
|
||||
def get_pexels_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
|
||||
all_images = []
|
||||
next_page = f"https://api.pexels.com/v1/collections/{config.collection_id}?page=1&per_page=80&type=photos"
|
||||
|
||||
while True:
|
||||
response = requests.get(next_page, headers={
|
||||
"Authorization": f"{config.api_key}"
|
||||
})
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
all_images.extend(data['media'])
|
||||
if 'next_page' in data and data['next_page']:
|
||||
next_page = data['next_page']
|
||||
else:
|
||||
break
|
||||
|
||||
photos = []
|
||||
for image in all_images:
|
||||
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
|
||||
url = f"{image['src']['original']}?auto=compress&cs=tinysrgb&h={new_height}&w={new_width}"
|
||||
filename = os.path.basename(image['src']['original'])
|
||||
|
||||
photos.append(Photo(
|
||||
id=image['id'],
|
||||
host="pexels",
|
||||
width=image['width'],
|
||||
height=image['height'],
|
||||
url=url,
|
||||
filename=filename
|
||||
))
|
||||
|
||||
return photos
|
||||
|
||||
|
||||
def get_unsplash_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
|
||||
headers = {
|
||||
# "Authorization": f"Client-ID {UNSPLASH_ACCESS_KEY}"
|
||||
"Authorization": f"Client-ID {config.api_key}"
|
||||
}
|
||||
# headers['Authorization'] = f"Bearer {token}"
|
||||
|
||||
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page=1&per_page=30"
|
||||
response = requests.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
res_headers = response.headers
|
||||
# parse the link header to get the next page
|
||||
# 'Link': '<https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=82>; rel="last", <https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=2>; rel="next"'
|
||||
has_next_page = False
|
||||
if 'Link' in res_headers:
|
||||
has_next_page = True
|
||||
link_header = res_headers['Link']
|
||||
link_header = link_header.split(',')
|
||||
link_header = [link.strip() for link in link_header]
|
||||
link_header = [link.split(';') for link in link_header]
|
||||
link_header = [[link[0].strip('<>'), link[1].strip().strip('"')] for link in link_header]
|
||||
link_header = {link[1]: link[0] for link in link_header}
|
||||
|
||||
# get page number from last url
|
||||
last_page = link_header['rel="last']
|
||||
last_page = last_page.split('?')[1]
|
||||
last_page = last_page.split('&')
|
||||
last_page = [param.split('=') for param in last_page]
|
||||
last_page = {param[0]: param[1] for param in last_page}
|
||||
last_page = int(last_page['page'])
|
||||
|
||||
all_images = response.json()
|
||||
|
||||
if has_next_page:
|
||||
# assume we start on page 1, so we don't need to get it again
|
||||
for page in tqdm.tqdm(range(2, last_page + 1)):
|
||||
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page={page}&per_page=30"
|
||||
response = requests.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
all_images.extend(response.json())
|
||||
|
||||
photos = []
|
||||
for image in all_images:
|
||||
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
|
||||
url = f"{image['urls']['raw']}&w={new_width}"
|
||||
filename = f"{image['id']}.jpg"
|
||||
|
||||
photos.append(Photo(
|
||||
id=image['id'],
|
||||
host="unsplash",
|
||||
width=image['width'],
|
||||
height=image['height'],
|
||||
url=url,
|
||||
filename=filename
|
||||
))
|
||||
|
||||
return photos
|
||||
|
||||
|
||||
def get_img_paths(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = os.listdir(dir_path)
|
||||
# remove non image files
|
||||
local_files = [file for file in local_files if os.path.splitext(file)[1].lower() in img_exts]
|
||||
# make full path
|
||||
local_files = [os.path.join(dir_path, file) for file in local_files]
|
||||
return local_files
|
||||
|
||||
|
||||
def get_local_image_ids(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = get_img_paths(dir_path)
|
||||
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
|
||||
return set([os.path.basename(file).split('.')[0] for file in local_files])
|
||||
|
||||
|
||||
def get_local_image_file_names(dir_path: str):
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
local_files = get_img_paths(dir_path)
|
||||
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
|
||||
return set([os.path.basename(file) for file in local_files])
|
||||
|
||||
|
||||
def download_image(photo: Photo, dir_path: str, min_width: int = 1024, min_height: int = 1024):
|
||||
img_width = photo.width
|
||||
img_height = photo.height
|
||||
|
||||
if img_width < min_width or img_height < min_height:
|
||||
raise ValueError(f"Skipping {photo.id} because it is too small: {img_width}x{img_height}")
|
||||
|
||||
img_response = requests.get(photo.url)
|
||||
img_response.raise_for_status()
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
|
||||
filename = os.path.join(dir_path, photo.filename)
|
||||
with open(filename, 'wb') as file:
|
||||
file.write(img_response.content)
|
||||
|
||||
|
||||
def update_caption(img_path: str):
|
||||
# if the caption is a txt file, convert it to a json file
|
||||
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
# see if it exists
|
||||
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json")):
|
||||
# todo add poi and what not
|
||||
return # we have a json file
|
||||
caption = ""
|
||||
# see if txt file exists
|
||||
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt")):
|
||||
# read it
|
||||
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"), 'r') as file:
|
||||
caption = file.read()
|
||||
# write json file
|
||||
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json"), 'w') as file:
|
||||
file.write(f'{{"caption": "{caption}"}}')
|
||||
|
||||
# delete txt file
|
||||
os.remove(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"))
|
||||
|
||||
|
||||
# def equalize_img(img_path: str):
|
||||
# input_path = img_path
|
||||
# output_path = os.path.join(img_root_path(img_path), COLOR_CORRECTED_DIR, os.path.basename(img_path))
|
||||
# os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
# process_img(
|
||||
# img_path=input_path,
|
||||
# output_path=output_path,
|
||||
# equalize=True,
|
||||
# max_size=2056,
|
||||
# white_balance=False,
|
||||
# gamma_correction=False,
|
||||
# strength=0.6,
|
||||
# )
|
||||
|
||||
|
||||
# def annotate_depth(img_path: str):
|
||||
# # make fake args
|
||||
# args = argparse.Namespace()
|
||||
# args.annotator = "midas"
|
||||
# args.res = 1024
|
||||
#
|
||||
# img = cv2.imread(img_path)
|
||||
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
#
|
||||
# output = annotate(img, args)
|
||||
#
|
||||
# output = output.astype('uint8')
|
||||
# output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR)
|
||||
#
|
||||
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
|
||||
# output_path = os.path.join(img_root_path(img_path), DEPTH_DIR, os.path.basename(img_path))
|
||||
#
|
||||
# cv2.imwrite(output_path, output)
|
||||
|
||||
|
||||
# def invert_depth(img_path: str):
|
||||
# img = cv2.imread(img_path)
|
||||
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
# # invert the colors
|
||||
# img = cv2.bitwise_not(img)
|
||||
#
|
||||
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
|
||||
# output_path = os.path.join(img_root_path(img_path), INVERTED_DEPTH_DIR, os.path.basename(img_path))
|
||||
# cv2.imwrite(output_path, img)
|
||||
|
||||
|
||||
#
|
||||
# # update our list of raw images
|
||||
# raw_images = get_img_paths(raw_dir)
|
||||
#
|
||||
# # update raw captions
|
||||
# for image_id in tqdm.tqdm(raw_images, desc="Updating raw captions"):
|
||||
# update_caption(image_id)
|
||||
#
|
||||
# # equalize images
|
||||
# for img_path in tqdm.tqdm(raw_images, desc="Equalizing images"):
|
||||
# if img_path not in eq_images:
|
||||
# equalize_img(img_path)
|
||||
#
|
||||
# # update our list of eq images
|
||||
# eq_images = get_img_paths(eq_dir)
|
||||
# # update eq captions
|
||||
# for image_id in tqdm.tqdm(eq_images, desc="Updating eq captions"):
|
||||
# update_caption(image_id)
|
||||
#
|
||||
# # annotate depth
|
||||
# depth_dir = os.path.join(root_dir, DEPTH_DIR)
|
||||
# depth_images = get_img_paths(depth_dir)
|
||||
# for img_path in tqdm.tqdm(eq_images, desc="Annotating depth"):
|
||||
# if img_path not in depth_images:
|
||||
# annotate_depth(img_path)
|
||||
#
|
||||
# depth_images = get_img_paths(depth_dir)
|
||||
#
|
||||
# # invert depth
|
||||
# inv_depth_dir = os.path.join(root_dir, INVERTED_DEPTH_DIR)
|
||||
# inv_depth_images = get_img_paths(inv_depth_dir)
|
||||
# for img_path in tqdm.tqdm(depth_images, desc="Inverting depth"):
|
||||
# if img_path not in inv_depth_images:
|
||||
# invert_depth(img_path)
|
||||
@@ -0,0 +1,235 @@
|
||||
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.config_modules import ReferenceDatasetConfig
|
||||
from toolkit.data_loader import PairedImageDataset
|
||||
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
import random
|
||||
from toolkit.basic import value_map
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class ReferenceSliderConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
|
||||
self.weight_jitter: float = kwargs.get('weight_jitter', 0.0)
|
||||
self.datasets: List[ReferenceDatasetConfig] = [ReferenceDatasetConfig(**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,
|
||||
'pos_weight': dataset.pos_weight,
|
||||
'neg_weight': dataset.neg_weight,
|
||||
'pos_folder': dataset.pos_folder,
|
||||
'neg_folder': dataset.neg_folder,
|
||||
}
|
||||
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):
|
||||
with torch.no_grad():
|
||||
imgs, prompts, network_weights = batch
|
||||
network_pos_weight, network_neg_weight = network_weights
|
||||
|
||||
if isinstance(network_pos_weight, torch.Tensor):
|
||||
network_pos_weight = network_pos_weight.item()
|
||||
if isinstance(network_neg_weight, torch.Tensor):
|
||||
network_neg_weight = network_neg_weight.item()
|
||||
|
||||
# get an array of random floats between -weight_jitter and weight_jitter
|
||||
loss_jitter_multiplier = 1.0
|
||||
weight_jitter = self.slider_config.weight_jitter
|
||||
if weight_jitter > 0.0:
|
||||
jitter_list = random.uniform(-weight_jitter, weight_jitter)
|
||||
orig_network_pos_weight = network_pos_weight
|
||||
network_pos_weight += jitter_list
|
||||
network_neg_weight += (jitter_list * -1.0)
|
||||
# penalize the loss for its distance from network_pos_weight
|
||||
# a jitter_list of abs(3.0) on a weight of 5.0 is a 60% jitter
|
||||
# so the loss_jitter_multiplier needs to be 0.4
|
||||
loss_jitter_multiplier = value_map(abs(jitter_list), 0.0, weight_jitter, 1.0, 0.0)
|
||||
|
||||
|
||||
# if items in network_weight list are tensors, convert them to floats
|
||||
|
||||
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)
|
||||
|
||||
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 = [network_pos_weight * 1.0, network_neg_weight * -1.0]
|
||||
|
||||
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):
|
||||
# fix issue with them being tuples sometimes
|
||||
prompt_list = []
|
||||
for prompt in prompts:
|
||||
if isinstance(prompt, tuple):
|
||||
prompt = prompt[0]
|
||||
prompt_list.append(prompt)
|
||||
conditional_embeds = self.sd.encode_prompt(prompt_list).to(self.device_torch, dtype=dtype)
|
||||
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
|
||||
# if self.model_config.is_xl:
|
||||
# # todo also allow for setting this for low ram in general, but sdxl spikes a ton on back prop
|
||||
# network_multiplier_list = network_multiplier
|
||||
# noisy_latent_list = torch.chunk(noisy_latents, 2, dim=0)
|
||||
# noise_list = torch.chunk(noise, 2, dim=0)
|
||||
# timesteps_list = torch.chunk(timesteps, 2, dim=0)
|
||||
# conditional_embeds_list = split_prompt_embeds(conditional_embeds)
|
||||
# else:
|
||||
network_multiplier_list = [network_multiplier]
|
||||
noisy_latent_list = [noisy_latents]
|
||||
noise_list = [noise]
|
||||
timesteps_list = [timesteps]
|
||||
conditional_embeds_list = [conditional_embeds]
|
||||
|
||||
losses = []
|
||||
# allow to chunk it out to save vram
|
||||
for network_multiplier, noisy_latents, noise, timesteps, conditional_embeds in zip(
|
||||
network_multiplier_list, noisy_latent_list, noise_list, timesteps_list, conditional_embeds_list
|
||||
):
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
|
||||
self.network.multiplier = network_multiplier
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
)
|
||||
noise = noise.to(self.device_torch, dtype=dtype)
|
||||
|
||||
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])
|
||||
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, noise_scheduler, self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean() * loss_jitter_multiplier
|
||||
|
||||
loss_float = loss.item()
|
||||
losses.append(loss_float)
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
|
||||
# apply gradients
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': sum(losses) / len(losses) if len(losses) > 0 else 0.0}
|
||||
)
|
||||
|
||||
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.
|
||||
640
extensions_built_in/sd_trainer/SDTrainer.py
Normal file
640
extensions_built_in/sd_trainer/SDTrainer.py
Normal file
@@ -0,0 +1,640 @@
|
||||
from collections import OrderedDict
|
||||
from typing import Union, Literal, List
|
||||
from diffusers import T2IAdapter
|
||||
|
||||
from toolkit import train_tools
|
||||
from toolkit.basic import value_map, adain, get_mean_std
|
||||
from toolkit.config_modules import GuidanceConfig
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO, FileItemDTO
|
||||
from toolkit.image_utils import show_tensors, show_latents
|
||||
from toolkit.ip_adapter import IPAdapter
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, BlankNetwork
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight, add_all_snr_to_noise_scheduler, \
|
||||
apply_learnable_snr_gos, LearnableSNRGamma
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
adapter_transforms = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
|
||||
|
||||
class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
self.assistant_adapter: Union['T2IAdapter', None]
|
||||
self.do_prior_prediction = False
|
||||
self.do_long_prompts = False
|
||||
if self.train_config.inverted_mask_prior:
|
||||
self.do_prior_prediction = True
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def before_dataset_load(self):
|
||||
self.assistant_adapter = None
|
||||
# get adapter assistant if one is set
|
||||
if self.train_config.adapter_assist_name_or_path is not None:
|
||||
adapter_path = self.train_config.adapter_assist_name_or_path
|
||||
|
||||
# dont name this adapter since we are not training it
|
||||
self.assistant_adapter = T2IAdapter.from_pretrained(
|
||||
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype), varient="fp16"
|
||||
).to(self.device_torch)
|
||||
self.assistant_adapter.eval()
|
||||
self.assistant_adapter.requires_grad_(False)
|
||||
flush()
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
# move vae to device if we did not cache latents
|
||||
if not self.is_latents_cached:
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
else:
|
||||
# offload it. Already cached
|
||||
self.sd.vae.to('cpu')
|
||||
flush()
|
||||
add_all_snr_to_noise_scheduler(self.sd.noise_scheduler, self.device_torch)
|
||||
|
||||
# you can expand these in a child class to make customization easier
|
||||
def calculate_loss(
|
||||
self,
|
||||
noise_pred: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
noisy_latents: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
mask_multiplier: Union[torch.Tensor, float] = 1.0,
|
||||
prior_pred: Union[torch.Tensor, None] = None,
|
||||
**kwargs
|
||||
):
|
||||
loss_target = self.train_config.loss_target
|
||||
|
||||
prior_mask_multiplier = None
|
||||
target_mask_multiplier = None
|
||||
|
||||
if self.train_config.match_noise_norm:
|
||||
# match the norm of the noise
|
||||
noise_norm = torch.linalg.vector_norm(noise, ord=2, dim=(1, 2, 3), keepdim=True)
|
||||
noise_pred_norm = torch.linalg.vector_norm(noise_pred, ord=2, dim=(1, 2, 3), keepdim=True)
|
||||
noise_pred = noise_pred * (noise_norm / noise_pred_norm)
|
||||
|
||||
if self.train_config.inverted_mask_prior:
|
||||
# we need to make the noise prediction be a masked blending of noise and prior_pred
|
||||
prior_mask_multiplier = 1.0 - mask_multiplier
|
||||
# target_mask_multiplier = mask_multiplier
|
||||
# mask_multiplier = 1.0
|
||||
target = noise
|
||||
# target = (noise * mask_multiplier) + (prior_pred * prior_mask_multiplier)
|
||||
# set masked multiplier to 1.0 so we dont double apply it
|
||||
# mask_multiplier = 1.0
|
||||
elif prior_pred is not None:
|
||||
# matching adapter prediction
|
||||
target = prior_pred
|
||||
elif self.sd.prediction_type == 'v_prediction':
|
||||
# v-parameterization training
|
||||
target = self.sd.noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
|
||||
else:
|
||||
target = noise
|
||||
|
||||
pred = noise_pred
|
||||
|
||||
ignore_snr = False
|
||||
|
||||
if loss_target == 'source' or loss_target == 'unaugmented':
|
||||
# ignore_snr = True
|
||||
if batch.sigmas is None:
|
||||
raise ValueError("Batch sigmas is None. This should not happen")
|
||||
|
||||
# src https://github.com/huggingface/diffusers/blob/324d18fba23f6c9d7475b0ff7c777685f7128d40/examples/t2i_adapter/train_t2i_adapter_sdxl.py#L1190
|
||||
denoised_latents = noise_pred * (-batch.sigmas) + noisy_latents
|
||||
weighing = batch.sigmas ** -2.0
|
||||
if loss_target == 'source':
|
||||
# denoise the latent and compare to the latent in the batch
|
||||
target = batch.latents
|
||||
elif loss_target == 'unaugmented':
|
||||
# we have to encode images into latents for now
|
||||
# we also denoise as the unaugmented tensor is not a noisy diffirental
|
||||
with torch.no_grad():
|
||||
unaugmented_latents = self.sd.encode_images(batch.unaugmented_tensor)
|
||||
unaugmented_latents = unaugmented_latents * self.train_config.latent_multiplier
|
||||
target = unaugmented_latents.detach()
|
||||
|
||||
# Get the target for loss depending on the prediction type
|
||||
if self.sd.noise_scheduler.config.prediction_type == "epsilon":
|
||||
target = target # we are computing loss against denoise latents
|
||||
elif self.sd.noise_scheduler.config.prediction_type == "v_prediction":
|
||||
target = self.sd.noise_scheduler.get_velocity(target, noise, timesteps)
|
||||
else:
|
||||
raise ValueError(f"Unknown prediction type {self.sd.noise_scheduler.config.prediction_type}")
|
||||
|
||||
# mse loss without reduction
|
||||
loss_per_element = (weighing.float() * (denoised_latents.float() - target.float()) ** 2)
|
||||
loss = loss_per_element
|
||||
else:
|
||||
loss = torch.nn.functional.mse_loss(pred.float(), target.float(), reduction="none")
|
||||
|
||||
# multiply by our mask
|
||||
loss = loss * mask_multiplier
|
||||
|
||||
if self.train_config.inverted_mask_prior:
|
||||
# to a loss to unmasked areas of the prior for unmasked regularization
|
||||
prior_loss = torch.nn.functional.mse_loss(
|
||||
prior_pred.float(),
|
||||
pred.float(),
|
||||
reduction="none"
|
||||
)
|
||||
prior_loss = prior_loss * prior_mask_multiplier * self.train_config.inverted_mask_prior_multiplier
|
||||
loss = loss + prior_loss
|
||||
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
if self.train_config.learnable_snr_gos:
|
||||
# add snr_gamma
|
||||
loss = apply_learnable_snr_gos(loss, timesteps, self.snr_gos)
|
||||
elif self.train_config.snr_gamma is not None and self.train_config.snr_gamma > 0.000001 and not ignore_snr:
|
||||
# add snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler, self.train_config.snr_gamma, fixed=True)
|
||||
elif self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001 and not ignore_snr:
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler, self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean()
|
||||
return loss
|
||||
|
||||
def preprocess_batch(self, batch: 'DataLoaderBatchDTO'):
|
||||
return batch
|
||||
|
||||
def get_guided_loss(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
# Perform targeted guidance (working title)
|
||||
conditional_noisy_latents = noisy_latents # target images
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
if batch.unconditional_latents is not None:
|
||||
# unconditional latents are the "neutral" images. Add noise here identical to
|
||||
# the noise added to the conditional latents, at the same timesteps
|
||||
unconditional_noisy_latents = self.sd.noise_scheduler.add_noise(
|
||||
batch.unconditional_latents, noise, timesteps
|
||||
)
|
||||
|
||||
# calculate the differential between our conditional (target image) and out unconditional (neutral image)
|
||||
target_differential_noise = unconditional_noisy_latents - conditional_noisy_latents
|
||||
target_differential_noise = target_differential_noise.detach()
|
||||
|
||||
# add the target differential to the target latents as if it were noise with the scheduler, scaled to
|
||||
# the current timestep. Scaling the noise here is important as it scales our guidance to the current
|
||||
# timestep. This is the key to making the guidance work.
|
||||
guidance_latents = self.sd.noise_scheduler.add_noise(
|
||||
conditional_noisy_latents,
|
||||
target_differential_noise,
|
||||
timesteps
|
||||
)
|
||||
|
||||
# Disable the LoRA network so we can predict parent network knowledge without it
|
||||
self.network.is_active = False
|
||||
self.sd.unet.eval()
|
||||
|
||||
# Predict noise to get a baseline of what the parent network wants to do with the latents + noise.
|
||||
# This acts as our control to preserve the unaltered parts of the image.
|
||||
baseline_prediction = self.sd.predict_noise(
|
||||
latents=guidance_latents.to(self.device_torch, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
).detach()
|
||||
|
||||
# turn the LoRA network back on.
|
||||
self.sd.unet.train()
|
||||
self.network.is_active = True
|
||||
self.network.multiplier = network_weight_list
|
||||
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
prediction = self.sd.predict_noise(
|
||||
latents=guidance_latents.to(self.device_torch, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
# remove the baseline prediction from our prediction to get the differential between the two
|
||||
# all that should be left is the differential between the conditional and unconditional images
|
||||
pred_differential_noise = prediction - baseline_prediction
|
||||
|
||||
# for loss, we target ONLY the unscaled differential between our conditional and unconditional latents
|
||||
# not the timestep scaled noise that was added. This is the diffusion training process.
|
||||
# This will guide the network to make identical predictions it previously did for everything EXCEPT our
|
||||
# differential between the conditional and unconditional images (target)
|
||||
loss = torch.nn.functional.mse_loss(
|
||||
pred_differential_noise.float(),
|
||||
target_differential_noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
loss = loss.mean([1, 2, 3])
|
||||
loss = self.apply_snr(loss, timesteps)
|
||||
loss = loss.mean()
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
def get_prior_prediction(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
**kwargs
|
||||
):
|
||||
# do a prediction here so we can match its output with network multiplier set to 0.0
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
# dont use network on this
|
||||
# self.network.multiplier = 0.0
|
||||
was_network_active = self.network.is_active
|
||||
self.network.is_active = False
|
||||
self.sd.unet.eval()
|
||||
prior_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
self.sd.unet.train()
|
||||
prior_pred = prior_pred.detach()
|
||||
# remove the residuals as we wont use them on prediction when matching control
|
||||
if match_adapter_assist and 'down_block_additional_residuals' in pred_kwargs:
|
||||
del pred_kwargs['down_block_additional_residuals']
|
||||
# restore network
|
||||
# self.network.multiplier = network_weight_list
|
||||
self.network.is_active = was_network_active
|
||||
return prior_pred
|
||||
|
||||
def before_unet_predict(self):
|
||||
pass
|
||||
|
||||
def after_unet_predict(self):
|
||||
pass
|
||||
|
||||
def end_of_training_loop(self):
|
||||
pass
|
||||
|
||||
def hook_train_loop(self, batch: 'DataLoaderBatchDTO'):
|
||||
self.timer.start('preprocess_batch')
|
||||
batch = self.preprocess_batch(batch)
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
|
||||
network_weight_list = batch.get_network_weight_list()
|
||||
if self.train_config.single_item_batching:
|
||||
network_weight_list = network_weight_list + network_weight_list
|
||||
|
||||
has_adapter_img = batch.control_tensor is not None
|
||||
|
||||
match_adapter_assist = False
|
||||
|
||||
|
||||
# check if we are matching the adapter assistant
|
||||
if self.assistant_adapter:
|
||||
if self.train_config.match_adapter_chance == 1.0:
|
||||
match_adapter_assist = True
|
||||
elif self.train_config.match_adapter_chance > 0.0:
|
||||
match_adapter_assist = torch.rand(
|
||||
(1,), device=self.device_torch, dtype=dtype
|
||||
) < self.train_config.match_adapter_chance
|
||||
|
||||
self.timer.stop('preprocess_batch')
|
||||
|
||||
with torch.no_grad():
|
||||
loss_multiplier = torch.ones((noisy_latents.shape[0], 1, 1, 1), device=self.device_torch, dtype=dtype)
|
||||
for idx, file_item in enumerate(batch.file_items):
|
||||
if file_item.is_reg:
|
||||
loss_multiplier[idx] = loss_multiplier[idx] * self.train_config.reg_weight
|
||||
|
||||
|
||||
adapter_images = None
|
||||
sigmas = None
|
||||
if has_adapter_img and (self.adapter or self.assistant_adapter):
|
||||
with self.timer('get_adapter_images'):
|
||||
# todo move this to data loader
|
||||
if batch.control_tensor is not None:
|
||||
adapter_images = batch.control_tensor.to(self.device_torch, dtype=dtype).detach()
|
||||
# match in channels
|
||||
if self.assistant_adapter is not None:
|
||||
in_channels = self.assistant_adapter.config.in_channels
|
||||
if adapter_images.shape[1] != in_channels:
|
||||
# we need to match the channels
|
||||
adapter_images = adapter_images[:, :in_channels, :, :]
|
||||
else:
|
||||
raise NotImplementedError("Adapter images now must be loaded with dataloader")
|
||||
# not 100% sure what this does. But they do it here
|
||||
# https://github.com/huggingface/diffusers/blob/38a664a3d61e27ab18cd698231422b3c38d6eebf/examples/t2i_adapter/train_t2i_adapter_sdxl.py#L1170
|
||||
# sigmas = self.get_sigmas(timesteps, len(noisy_latents.shape), noisy_latents.dtype)
|
||||
# noisy_latents = noisy_latents / ((sigmas ** 2 + 1) ** 0.5)
|
||||
|
||||
mask_multiplier = torch.ones((noisy_latents.shape[0], 1, 1, 1), device=self.device_torch, dtype=dtype)
|
||||
if batch.mask_tensor is not None:
|
||||
with self.timer('get_mask_multiplier'):
|
||||
# upsampling no supported for bfloat16
|
||||
mask_multiplier = batch.mask_tensor.to(self.device_torch, dtype=torch.float16).detach()
|
||||
# scale down to the size of the latents, mask multiplier shape(bs, 1, width, height), noisy_latents shape(bs, channels, width, height)
|
||||
mask_multiplier = torch.nn.functional.interpolate(
|
||||
mask_multiplier, size=(noisy_latents.shape[2], noisy_latents.shape[3])
|
||||
)
|
||||
# expand to match latents
|
||||
mask_multiplier = mask_multiplier.expand(-1, noisy_latents.shape[1], -1, -1)
|
||||
mask_multiplier = mask_multiplier.to(self.device_torch, dtype=dtype).detach()
|
||||
|
||||
def get_adapter_multiplier():
|
||||
if self.adapter and isinstance(self.adapter, T2IAdapter):
|
||||
# training a t2i adapter, not using as assistant.
|
||||
return 1.0
|
||||
elif match_adapter_assist:
|
||||
# training a texture. We want it high
|
||||
adapter_strength_min = 0.9
|
||||
adapter_strength_max = 1.0
|
||||
else:
|
||||
# training with assistance, we want it low
|
||||
adapter_strength_min = 0.4
|
||||
adapter_strength_max = 0.7
|
||||
# adapter_strength_min = 0.9
|
||||
# adapter_strength_max = 1.1
|
||||
|
||||
adapter_conditioning_scale = torch.rand(
|
||||
(1,), device=self.device_torch, dtype=dtype
|
||||
)
|
||||
|
||||
adapter_conditioning_scale = value_map(
|
||||
adapter_conditioning_scale,
|
||||
0.0,
|
||||
1.0,
|
||||
adapter_strength_min,
|
||||
adapter_strength_max
|
||||
)
|
||||
return adapter_conditioning_scale
|
||||
|
||||
# flush()
|
||||
with self.timer('grad_setup'):
|
||||
|
||||
# text encoding
|
||||
grad_on_text_encoder = False
|
||||
if self.train_config.train_text_encoder:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
if self.embedding:
|
||||
grad_on_text_encoder = True
|
||||
|
||||
# have a blank network so we can wrap it in a context and set multipliers without checking every time
|
||||
if self.network is not None:
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
# set the weights
|
||||
network.multiplier = network_weight_list
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
# activate network if it exits
|
||||
|
||||
prompts_1 = conditioned_prompts
|
||||
prompts_2 = None
|
||||
if self.train_config.short_and_long_captions_encoder_split and self.sd.is_xl:
|
||||
prompts_1 = batch.get_caption_short_list()
|
||||
prompts_2 = conditioned_prompts
|
||||
|
||||
# make the batch splits
|
||||
if self.train_config.single_item_batching:
|
||||
if self.model_config.refiner_name_or_path is not None:
|
||||
raise ValueError("Single item batching is not supported when training the refiner")
|
||||
batch_size = noisy_latents.shape[0]
|
||||
# chunk/split everything
|
||||
noisy_latents_list = torch.chunk(noisy_latents, batch_size, dim=0)
|
||||
noise_list = torch.chunk(noise, batch_size, dim=0)
|
||||
timesteps_list = torch.chunk(timesteps, batch_size, dim=0)
|
||||
conditioned_prompts_list = [[prompt] for prompt in prompts_1]
|
||||
if imgs is not None:
|
||||
imgs_list = torch.chunk(imgs, batch_size, dim=0)
|
||||
else:
|
||||
imgs_list = [None for _ in range(batch_size)]
|
||||
if adapter_images is not None:
|
||||
adapter_images_list = torch.chunk(adapter_images, batch_size, dim=0)
|
||||
else:
|
||||
adapter_images_list = [None for _ in range(batch_size)]
|
||||
mask_multiplier_list = torch.chunk(mask_multiplier, batch_size, dim=0)
|
||||
if prompts_2 is None:
|
||||
prompt_2_list = [None for _ in range(batch_size)]
|
||||
else:
|
||||
prompt_2_list = [[prompt] for prompt in prompts_2]
|
||||
|
||||
else:
|
||||
noisy_latents_list = [noisy_latents]
|
||||
noise_list = [noise]
|
||||
timesteps_list = [timesteps]
|
||||
conditioned_prompts_list = [prompts_1]
|
||||
imgs_list = [imgs]
|
||||
adapter_images_list = [adapter_images]
|
||||
mask_multiplier_list = [mask_multiplier]
|
||||
if prompts_2 is None:
|
||||
prompt_2_list = [None]
|
||||
else:
|
||||
prompt_2_list = [prompts_2]
|
||||
|
||||
for noisy_latents, noise, timesteps, conditioned_prompts, imgs, adapter_images, mask_multiplier, prompt_2 in zip(
|
||||
noisy_latents_list,
|
||||
noise_list,
|
||||
timesteps_list,
|
||||
conditioned_prompts_list,
|
||||
imgs_list,
|
||||
adapter_images_list,
|
||||
mask_multiplier_list,
|
||||
prompt_2_list
|
||||
):
|
||||
if self.train_config.negative_prompt is not None:
|
||||
# add negative prompt
|
||||
conditioned_prompts = conditioned_prompts + [self.train_config.negative_prompt for x in
|
||||
range(len(conditioned_prompts))]
|
||||
if prompt_2 is not None:
|
||||
prompt_2 = prompt_2 + [self.train_config.negative_prompt for x in range(len(prompt_2))]
|
||||
|
||||
with network:
|
||||
with self.timer('encode_prompt'):
|
||||
if grad_on_text_encoder:
|
||||
with torch.set_grad_enabled(True):
|
||||
conditional_embeds = self.sd.encode_prompt(
|
||||
conditioned_prompts, prompt_2,
|
||||
dropout_prob=self.train_config.prompt_dropout_prob,
|
||||
long_prompts=self.do_long_prompts).to(
|
||||
self.device_torch,
|
||||
dtype=dtype)
|
||||
else:
|
||||
with torch.set_grad_enabled(False):
|
||||
# make sure it is in eval mode
|
||||
if isinstance(self.sd.text_encoder, list):
|
||||
for te in self.sd.text_encoder:
|
||||
te.eval()
|
||||
else:
|
||||
self.sd.text_encoder.eval()
|
||||
conditional_embeds = self.sd.encode_prompt(
|
||||
conditioned_prompts, prompt_2,
|
||||
dropout_prob=self.train_config.prompt_dropout_prob,
|
||||
long_prompts=self.do_long_prompts).to(
|
||||
self.device_torch,
|
||||
dtype=dtype)
|
||||
|
||||
# detach the embeddings
|
||||
conditional_embeds = conditional_embeds.detach()
|
||||
|
||||
# flush()
|
||||
pred_kwargs = {}
|
||||
if has_adapter_img and (
|
||||
(self.adapter and isinstance(self.adapter, T2IAdapter)) or self.assistant_adapter):
|
||||
with torch.set_grad_enabled(self.adapter is not None):
|
||||
adapter = self.adapter if self.adapter else self.assistant_adapter
|
||||
adapter_multiplier = get_adapter_multiplier()
|
||||
with self.timer('encode_adapter'):
|
||||
down_block_additional_residuals = adapter(adapter_images)
|
||||
if self.assistant_adapter:
|
||||
# not training. detach
|
||||
down_block_additional_residuals = [
|
||||
sample.to(dtype=dtype).detach() * adapter_multiplier for sample in
|
||||
down_block_additional_residuals
|
||||
]
|
||||
else:
|
||||
down_block_additional_residuals = [
|
||||
sample.to(dtype=dtype) * adapter_multiplier for sample in
|
||||
down_block_additional_residuals
|
||||
]
|
||||
|
||||
pred_kwargs['down_block_additional_residuals'] = down_block_additional_residuals
|
||||
|
||||
prior_pred = None
|
||||
if (has_adapter_img and self.assistant_adapter and match_adapter_assist) or self.do_prior_prediction:
|
||||
with self.timer('prior predict'):
|
||||
prior_pred = self.get_prior_prediction(
|
||||
noisy_latents=noisy_latents,
|
||||
conditional_embeds=conditional_embeds,
|
||||
match_adapter_assist=match_adapter_assist,
|
||||
network_weight_list=network_weight_list,
|
||||
timesteps=timesteps,
|
||||
pred_kwargs=pred_kwargs,
|
||||
noise=noise,
|
||||
batch=batch,
|
||||
)
|
||||
|
||||
if has_adapter_img and self.adapter and isinstance(self.adapter, IPAdapter):
|
||||
with self.timer('encode_adapter'):
|
||||
with torch.no_grad():
|
||||
conditional_clip_embeds = self.adapter.get_clip_image_embeds_from_tensors(adapter_images)
|
||||
conditional_embeds = self.adapter(conditional_embeds, conditional_clip_embeds)
|
||||
|
||||
self.before_unet_predict()
|
||||
# do a prior pred if we have an unconditional image, we will swap out the giadance later
|
||||
if batch.unconditional_latents is not None:
|
||||
# do guided loss
|
||||
loss = self.get_guided_loss(
|
||||
noisy_latents=noisy_latents,
|
||||
conditional_embeds=conditional_embeds,
|
||||
match_adapter_assist=match_adapter_assist,
|
||||
network_weight_list=network_weight_list,
|
||||
timesteps=timesteps,
|
||||
pred_kwargs=pred_kwargs,
|
||||
batch=batch,
|
||||
noise=noise,
|
||||
)
|
||||
|
||||
else:
|
||||
with self.timer('predict_unet'):
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs
|
||||
)
|
||||
self.after_unet_predict()
|
||||
|
||||
with self.timer('calculate_loss'):
|
||||
noise = noise.to(self.device_torch, dtype=dtype).detach()
|
||||
loss = self.calculate_loss(
|
||||
noise_pred=noise_pred,
|
||||
noise=noise,
|
||||
noisy_latents=noisy_latents,
|
||||
timesteps=timesteps,
|
||||
batch=batch,
|
||||
mask_multiplier=mask_multiplier,
|
||||
prior_pred=prior_pred,
|
||||
)
|
||||
# check if nan
|
||||
if torch.isnan(loss):
|
||||
raise ValueError("loss is nan")
|
||||
|
||||
with self.timer('backward'):
|
||||
# todo we have multiplier seperated. works for now as res are not in same batch, but need to change
|
||||
loss = loss * loss_multiplier.mean()
|
||||
# IMPORTANT if gradient checkpointing do not leave with network when doing backward
|
||||
# it will destroy the gradients. This is because the network is a context manager
|
||||
# and will change the multipliers back to 0.0 when exiting. They will be
|
||||
# 0.0 for the backward pass and the gradients will be 0.0
|
||||
# I spent weeks on fighting this. DON'T DO IT
|
||||
# with fsdp_overlap_step_with_backward():
|
||||
loss.backward()
|
||||
# flush()
|
||||
|
||||
if not self.is_grad_accumulation_step:
|
||||
torch.nn.utils.clip_grad_norm_(self.params, self.train_config.max_grad_norm)
|
||||
# only step if we are not accumulating
|
||||
with self.timer('optimizer_step'):
|
||||
# apply gradients
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
else:
|
||||
# gradient accumulation. Just a place for breakpoint
|
||||
pass
|
||||
|
||||
# TODO Should we only step scheduler on grad step? If so, need to recalculate last step
|
||||
with self.timer('scheduler_step'):
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.embedding is not None:
|
||||
with self.timer('restore_embeddings'):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
self.embedding.restore_embeddings()
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss.item()}
|
||||
)
|
||||
|
||||
self.end_of_training_loop()
|
||||
|
||||
return loss_dict
|
||||
30
extensions_built_in/sd_trainer/__init__.py
Normal file
30
extensions_built_in/sd_trainer/__init__.py
Normal file
@@ -0,0 +1,30 @@
|
||||
# This is an example extension for custom training. It is great for experimenting with new ideas.
|
||||
from toolkit.extension import Extension
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class SDTrainerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "sd_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "SD 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 .SDTrainer import SDTrainer
|
||||
return SDTrainer
|
||||
|
||||
|
||||
# for backwards compatability
|
||||
class TextualInversionTrainer(SDTrainerExtension):
|
||||
uid = "textual_inversion_trainer"
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
SDTrainerExtension, TextualInversionTrainer
|
||||
]
|
||||
91
extensions_built_in/sd_trainer/config/train.example.yaml
Normal file
91
extensions_built_in/sd_trainer/config/train.example.yaml
Normal file
@@ -0,0 +1,91 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
name: test_v1
|
||||
process:
|
||||
- type: 'textual_inversion_trainer'
|
||||
training_folder: "out/TI"
|
||||
device: cuda:0
|
||||
# for tensorboard logging
|
||||
log_dir: "out/.tensorboard"
|
||||
embedding:
|
||||
trigger: "your_trigger_here"
|
||||
tokens: 12
|
||||
init_words: "man with short brown hair"
|
||||
save_format: "safetensors" # 'safetensors' or 'pt'
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 100 # save every this many steps
|
||||
max_step_saves_to_keep: 5 # only affects step counts
|
||||
datasets:
|
||||
- folder_path: "/path/to/dataset"
|
||||
caption_ext: "txt"
|
||||
default_caption: "[trigger]"
|
||||
buckets: true
|
||||
resolution: 512
|
||||
train:
|
||||
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
|
||||
steps: 3000
|
||||
weight_jitter: 0.0
|
||||
lr: 5e-5
|
||||
train_unet: false
|
||||
gradient_checkpointing: true
|
||||
train_text_encoder: false
|
||||
optimizer: "adamw"
|
||||
# optimizer: "prodigy"
|
||||
optimizer_params:
|
||||
weight_decay: 1e-2
|
||||
lr_scheduler: "constant"
|
||||
max_denoising_steps: 1000
|
||||
batch_size: 4
|
||||
dtype: bf16
|
||||
xformers: true
|
||||
min_snr_gamma: 5.0
|
||||
# skip_first_sample: true
|
||||
noise_offset: 0.0 # not needed for this
|
||||
model:
|
||||
# objective reality v2
|
||||
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
|
||||
is_v2: false # for v2 models
|
||||
is_xl: false # for SDXL models
|
||||
is_v_pred: false # for v-prediction models (most v2 models)
|
||||
sample:
|
||||
sampler: "ddpm" # must match train.noise_scheduler
|
||||
sample_every: 100 # sample every this many steps
|
||||
width: 512
|
||||
height: 512
|
||||
prompts:
|
||||
- "photo of [trigger] laughing"
|
||||
- "photo of [trigger] smiling"
|
||||
- "[trigger] close up"
|
||||
- "dark scene [trigger] frozen"
|
||||
- "[trigger] nighttime"
|
||||
- "a painting of [trigger]"
|
||||
- "a drawing of [trigger]"
|
||||
- "a cartoon of [trigger]"
|
||||
- "[trigger] pixar style"
|
||||
- "[trigger] costume"
|
||||
neg: ""
|
||||
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
|
||||
|
||||
# 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.
|
||||
# It is saved in the model so be aware of that. The software will include this
|
||||
# plus some other information for you automatically
|
||||
meta:
|
||||
# [name] gets replaced with the name above
|
||||
name: "[name]"
|
||||
# version: '1.0'
|
||||
# creator:
|
||||
# name: Your Name
|
||||
# email: your@gmail.com
|
||||
# website: https://your.website
|
||||
@@ -0,0 +1,533 @@
|
||||
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.config_modules import ReferenceDatasetConfig
|
||||
from toolkit.data_loader import PairedImageDataset
|
||||
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds, build_latent_image_batch_for_prompt_pair
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
|
||||
import gc
|
||||
from toolkit import train_tools
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
import random
|
||||
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
|
||||
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
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class UltimateSliderConfig(SliderConfig):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
|
||||
self.weight_jitter: float = kwargs.get('weight_jitter', 0.0)
|
||||
self.img_loss_weight: float = kwargs.get('img_loss_weight', 1.0)
|
||||
self.cfg_loss_weight: float = kwargs.get('cfg_loss_weight', 1.0)
|
||||
self.datasets: List[ReferenceDatasetConfig] = [ReferenceDatasetConfig(**d) for d in kwargs.get('datasets', [])]
|
||||
|
||||
|
||||
class UltimateSliderTrainerProcess(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 = UltimateSliderConfig(**self.get_conf('slider', {}))
|
||||
|
||||
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
|
||||
|
||||
# store a list of all the prompts from the dataset so we can cache it
|
||||
self.dataset_prompts = []
|
||||
self.train_with_dataset = self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0
|
||||
|
||||
def load_datasets(self):
|
||||
if self.data_loader is None and \
|
||||
self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0:
|
||||
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,
|
||||
'pos_weight': dataset.pos_weight,
|
||||
'neg_weight': dataset.neg_weight,
|
||||
'pos_folder': dataset.pos_folder,
|
||||
'neg_folder': dataset.neg_folder,
|
||||
}
|
||||
image_dataset = PairedImageDataset(config)
|
||||
datasets.append(image_dataset)
|
||||
|
||||
# capture all the prompts from it so we can cache the embeds
|
||||
self.dataset_prompts += image_dataset.get_all_prompts()
|
||||
|
||||
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):
|
||||
# load any datasets if they were passed
|
||||
self.load_datasets()
|
||||
|
||||
# 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()
|
||||
|
||||
# get encoded latents for our prompts
|
||||
with torch.no_grad():
|
||||
# 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 [""]
|
||||
|
||||
# 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
|
||||
|
||||
# remove duplicates
|
||||
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
|
||||
|
||||
# trim to max steps if max steps is lower than prompt count
|
||||
prompts_to_cache = prompts_to_cache[:self.train_config.steps]
|
||||
|
||||
if len(self.dataset_prompts) > 0:
|
||||
# add the prompts from the dataset
|
||||
prompts_to_cache += self.dataset_prompts
|
||||
|
||||
# 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]
|
||||
|
||||
# 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
|
||||
# end hook_before_train_loop
|
||||
|
||||
# move vae to device so we can encode on the fly
|
||||
# todo cache latents
|
||||
self.sd.vae.to(self.device_torch)
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.requires_grad_(False)
|
||||
|
||||
if self.train_config.gradient_checkpointing:
|
||||
# may get disabled elsewhere
|
||||
self.sd.unet.enable_gradient_checkpointing()
|
||||
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self, batch):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
### LOOP SETUP ###
|
||||
noise_scheduler = self.sd.noise_scheduler
|
||||
optimizer = self.optimizer
|
||||
lr_scheduler = self.lr_scheduler
|
||||
|
||||
### TARGET_PROMPTS ###
|
||||
# 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)
|
||||
|
||||
### PREP REFERENCE IMAGES ###
|
||||
|
||||
imgs, prompts, network_weights = batch
|
||||
network_pos_weight, network_neg_weight = network_weights
|
||||
|
||||
if isinstance(network_pos_weight, torch.Tensor):
|
||||
network_pos_weight = network_pos_weight.item()
|
||||
if isinstance(network_neg_weight, torch.Tensor):
|
||||
network_neg_weight = network_neg_weight.item()
|
||||
|
||||
# get an array of random floats between -weight_jitter and weight_jitter
|
||||
weight_jitter = self.slider_config.weight_jitter
|
||||
if weight_jitter > 0.0:
|
||||
jitter_list = random.uniform(-weight_jitter, weight_jitter)
|
||||
network_pos_weight += jitter_list
|
||||
network_neg_weight += (jitter_list * -1.0)
|
||||
|
||||
# if items in network_weight list are tensors, convert them to floats
|
||||
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)
|
||||
|
||||
height = positive_images.shape[2]
|
||||
width = positive_images.shape[3]
|
||||
batch_size = positive_images.shape[0]
|
||||
|
||||
positive_latents = self.sd.encode_images(positive_images)
|
||||
negative_latents = self.sd.encode_images(negative_images)
|
||||
|
||||
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)
|
||||
current_timestep_index = timesteps.item()
|
||||
current_timestep = noise_scheduler.timesteps[current_timestep_index]
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
### CFG SLIDER TRAINING PREP ###
|
||||
|
||||
# get CFG txt latents
|
||||
noisy_cfg_latents = build_latent_image_batch_for_prompt_pair(
|
||||
pos_latent=noisy_positive_latents,
|
||||
neg_latent=noisy_negative_latents,
|
||||
prompt_pair=prompt_pair,
|
||||
prompt_chunk_size=self.prompt_chunk_size,
|
||||
)
|
||||
noisy_cfg_latents.requires_grad = False
|
||||
|
||||
assert not self.network.is_active
|
||||
|
||||
# 4.20 GB RAM for 512x512
|
||||
positive_latents = self.sd.predict_noise(
|
||||
latents=noisy_cfg_latents,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.negative_target, # positive prompt
|
||||
self.train_config.batch_size,
|
||||
),
|
||||
timestep=current_timestep,
|
||||
guidance_scale=1.0
|
||||
)
|
||||
positive_latents.requires_grad = False
|
||||
|
||||
neutral_latents = self.sd.predict_noise(
|
||||
latents=noisy_cfg_latents,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.empty_prompt, # positive prompt (normally neutral
|
||||
self.train_config.batch_size,
|
||||
),
|
||||
timestep=current_timestep,
|
||||
guidance_scale=1.0
|
||||
)
|
||||
neutral_latents.requires_grad = False
|
||||
|
||||
unconditional_latents = self.sd.predict_noise(
|
||||
latents=noisy_cfg_latents,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.positive_target, # positive prompt
|
||||
self.train_config.batch_size,
|
||||
),
|
||||
timestep=current_timestep,
|
||||
guidance_scale=1.0
|
||||
)
|
||||
unconditional_latents.requires_grad = False
|
||||
|
||||
positive_latents_chunks = torch.chunk(positive_latents, self.prompt_chunk_size, dim=0)
|
||||
neutral_latents_chunks = torch.chunk(neutral_latents, self.prompt_chunk_size, dim=0)
|
||||
unconditional_latents_chunks = torch.chunk(unconditional_latents, self.prompt_chunk_size, dim=0)
|
||||
prompt_pair_chunks = split_prompt_pairs(prompt_pair, self.prompt_chunk_size)
|
||||
noisy_cfg_latents_chunks = torch.chunk(noisy_cfg_latents, self.prompt_chunk_size, dim=0)
|
||||
assert len(prompt_pair_chunks) == len(noisy_cfg_latents_chunks)
|
||||
|
||||
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 = [network_pos_weight * 1.0, network_neg_weight * -1.0]
|
||||
|
||||
flush()
|
||||
|
||||
loss_float = None
|
||||
loss_mirror_float = None
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
noisy_latents.requires_grad = False
|
||||
|
||||
# TODO allow both processed to train text encoder, for now, we just to unet and cache all text encodes
|
||||
# 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])
|
||||
|
||||
if self.train_with_dataset:
|
||||
embedding_list = []
|
||||
with torch.set_grad_enabled(self.train_config.train_text_encoder):
|
||||
for prompt in prompts:
|
||||
# get embedding form cache
|
||||
embedding = self.prompt_cache[prompt]
|
||||
embedding = embedding.to(self.device_torch, dtype=dtype)
|
||||
embedding_list.append(embedding)
|
||||
conditional_embeds = concat_prompt_embeds(embedding_list)
|
||||
# double up so we can do both sides of the slider
|
||||
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
else:
|
||||
# throw error. Not supported yet
|
||||
raise Exception("Datasets and targets required for ultimate slider")
|
||||
|
||||
if self.model_config.is_xl:
|
||||
# todo also allow for setting this for low ram in general, but sdxl spikes a ton on back prop
|
||||
network_multiplier_list = network_multiplier
|
||||
noisy_latent_list = torch.chunk(noisy_latents, 2, dim=0)
|
||||
noise_list = torch.chunk(noise, 2, dim=0)
|
||||
timesteps_list = torch.chunk(timesteps, 2, dim=0)
|
||||
conditional_embeds_list = split_prompt_embeds(conditional_embeds)
|
||||
else:
|
||||
network_multiplier_list = [network_multiplier]
|
||||
noisy_latent_list = [noisy_latents]
|
||||
noise_list = [noise]
|
||||
timesteps_list = [timesteps]
|
||||
conditional_embeds_list = [conditional_embeds]
|
||||
|
||||
## DO REFERENCE IMAGE TRAINING ##
|
||||
|
||||
reference_image_losses = []
|
||||
# allow to chunk it out to save vram
|
||||
for network_multiplier, noisy_latents, noise, timesteps, conditional_embeds in zip(
|
||||
network_multiplier_list, noisy_latent_list, noise_list, timesteps_list, conditional_embeds_list
|
||||
):
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
|
||||
self.network.multiplier = network_multiplier
|
||||
|
||||
noise_pred = self.sd.predict_noise(
|
||||
latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
|
||||
timestep=timesteps,
|
||||
)
|
||||
noise = noise.to(self.device_torch, dtype=dtype)
|
||||
|
||||
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
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps, noise_scheduler, self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean()
|
||||
loss = loss * self.slider_config.img_loss_weight
|
||||
loss_slide_float = loss.item()
|
||||
|
||||
loss_float = loss.item()
|
||||
reference_image_losses.append(loss_float)
|
||||
|
||||
# back propagate loss to free ram
|
||||
loss.backward()
|
||||
flush()
|
||||
|
||||
## DO CFG SLIDER TRAINING ##
|
||||
|
||||
cfg_loss_list = []
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
for prompt_pair_chunk, \
|
||||
noisy_cfg_latent_chunk, \
|
||||
positive_latents_chunk, \
|
||||
neutral_latents_chunk, \
|
||||
unconditional_latents_chunk \
|
||||
in zip(
|
||||
prompt_pair_chunks,
|
||||
noisy_cfg_latents_chunks,
|
||||
positive_latents_chunks,
|
||||
neutral_latents_chunks,
|
||||
unconditional_latents_chunks,
|
||||
):
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list
|
||||
|
||||
target_latents = self.sd.predict_noise(
|
||||
latents=noisy_cfg_latent_chunk,
|
||||
text_embeddings=train_tools.concat_prompt_embeddings(
|
||||
prompt_pair_chunk.positive_target, # negative prompt
|
||||
prompt_pair_chunk.target_class, # positive prompt
|
||||
self.train_config.batch_size,
|
||||
),
|
||||
timestep=current_timestep,
|
||||
guidance_scale=1.0
|
||||
)
|
||||
|
||||
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 = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
# match batch size
|
||||
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps_index_list, noise_scheduler,
|
||||
self.train_config.min_snr_gamma)
|
||||
|
||||
loss = loss.mean() * prompt_pair_chunk.weight * self.slider_config.cfg_loss_weight
|
||||
|
||||
loss.backward()
|
||||
cfg_loss_list.append(loss.item())
|
||||
del target_latents
|
||||
del offset_neutral
|
||||
del loss
|
||||
flush()
|
||||
|
||||
# apply gradients
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
|
||||
reference_image_loss = sum(reference_image_losses) / len(reference_image_losses) if len(
|
||||
reference_image_losses) > 0 else 0.0
|
||||
cfg_loss = sum(cfg_loss_list) / len(cfg_loss_list) if len(cfg_loss_list) > 0 else 0.0
|
||||
|
||||
loss_dict = OrderedDict({
|
||||
'loss/img': reference_image_loss,
|
||||
'loss/cfg': cfg_loss,
|
||||
})
|
||||
|
||||
return loss_dict
|
||||
# end hook_train_loop
|
||||
25
extensions_built_in/ultimate_slider_trainer/__init__.py
Normal file
25
extensions_built_in/ultimate_slider_trainer/__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 UltimateSliderTrainer(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "ultimate_slider_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Ultimate 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 .UltimateSliderTrainerProcess import UltimateSliderTrainerProcess
|
||||
return UltimateSliderTrainerProcess
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
UltimateSliderTrainer
|
||||
]
|
||||
@@ -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.1.0"
|
||||
|
||||
software_meta = v
|
||||
|
||||
@@ -6,19 +6,16 @@ from jobs.process import BaseProcess
|
||||
|
||||
|
||||
class BaseJob:
|
||||
config: OrderedDict
|
||||
job: str
|
||||
name: str
|
||||
meta: OrderedDict
|
||||
process: List[BaseProcess]
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
if not config:
|
||||
raise ValueError('config is required')
|
||||
self.process: List[BaseProcess]
|
||||
|
||||
self.config = config['config']
|
||||
self.raw_config = config
|
||||
self.job = config['job']
|
||||
self.torch_profiler = self.get_conf('torch_profiler', False)
|
||||
self.name = self.get_conf('name', required=True)
|
||||
if 'meta' in config:
|
||||
self.meta = config['meta']
|
||||
@@ -60,7 +57,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"]}')
|
||||
|
||||
22
jobs/ExtensionJob.py
Normal file
22
jobs/ExtensionJob.py
Normal file
@@ -0,0 +1,22 @@
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from jobs import BaseJob
|
||||
from toolkit.extension import get_all_extensions_process_dict
|
||||
from toolkit.paths import CONFIG_ROOT
|
||||
|
||||
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()
|
||||
31
jobs/GenerateJob.py
Normal file
31
jobs/GenerateJob.py
Normal file
@@ -0,0 +1,31 @@
|
||||
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):
|
||||
|
||||
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,13 +17,15 @@ sys.path.append(REPOS_ROOT)
|
||||
process_dict = {
|
||||
'vae': 'TrainVAEProcess',
|
||||
'slider': 'TrainSliderProcess',
|
||||
'slider_old': 'TrainSliderProcessOld',
|
||||
'lora_hack': 'TrainLoRAHack',
|
||||
'rescale_sd': 'TrainSDRescaleProcess',
|
||||
'esrgan': 'TrainESRGANProcess',
|
||||
'reference': 'TrainReferenceProcess',
|
||||
}
|
||||
|
||||
|
||||
class TrainJob(BaseJob):
|
||||
process: List[BaseExtractProcess]
|
||||
|
||||
def __init__(self, config: OrderedDict):
|
||||
super().__init__(config)
|
||||
@@ -34,18 +36,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 +47,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
|
||||
|
||||
19
jobs/process/BaseExtensionProcess.py
Normal file
19
jobs/process/BaseExtensionProcess.py
Normal file
@@ -0,0 +1,19 @@
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from jobs.process.BaseProcess import BaseProcess
|
||||
|
||||
|
||||
class BaseExtensionProcess(BaseProcess):
|
||||
def __init__(
|
||||
self,
|
||||
process_id: int,
|
||||
job,
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.process_id: int
|
||||
self.config: OrderedDict
|
||||
self.progress_bar: ForwardRef('tqdm') = None
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -12,11 +12,6 @@ from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
class BaseExtractProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
output_folder: str
|
||||
output_filename: str
|
||||
output_path: str
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -25,6 +20,10 @@ class BaseExtractProcess(BaseProcess):
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.config: OrderedDict
|
||||
self.output_folder: str
|
||||
self.output_filename: str
|
||||
self.output_path: str
|
||||
self.process_id = process_id
|
||||
self.job = job
|
||||
self.config = config
|
||||
|
||||
@@ -9,8 +9,6 @@ from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
|
||||
class BaseMergeProcess(BaseProcess):
|
||||
process_id: int
|
||||
config: OrderedDict
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -19,6 +17,8 @@ class BaseMergeProcess(BaseProcess):
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.process_id: int
|
||||
self.config: OrderedDict
|
||||
self.output_path = self.get_conf('output_path', required=True)
|
||||
self.dtype = self.get_conf('dtype', self.job.dtype)
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import copy
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
|
||||
from toolkit.timer import Timer
|
||||
|
||||
|
||||
class BaseProcess:
|
||||
meta: OrderedDict
|
||||
class BaseProcess(object):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -14,9 +14,15 @@ class BaseProcess:
|
||||
config: OrderedDict
|
||||
):
|
||||
self.process_id = process_id
|
||||
self.meta: OrderedDict
|
||||
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)
|
||||
self.timer: Timer = Timer(f'{self.name} Timer')
|
||||
self.performance_log_every = self.get_conf('performance_log_every', 0)
|
||||
|
||||
print(json.dumps(self.config, indent=4))
|
||||
|
||||
def get_conf(self, key, default=None, required=False, as_type=None):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,14 +1,21 @@
|
||||
import random
|
||||
from datetime import datetime
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef
|
||||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
import torch
|
||||
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
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -17,12 +24,30 @@ class BaseTrainProcess(BaseProcess):
|
||||
config: OrderedDict
|
||||
):
|
||||
super().__init__(process_id, job, config)
|
||||
self.process_id: int
|
||||
self.config: OrderedDict
|
||||
self.writer: 'SummaryWriter'
|
||||
self.job: Union['TrainJob', 'BaseJob', 'ExtensionJob']
|
||||
self.progress_bar: 'tqdm' = None
|
||||
|
||||
self.training_seed = self.get_conf('training_seed', self.job.training_seed if hasattr(self.job, 'training_seed') else None)
|
||||
# if training seed is set, use it
|
||||
if self.training_seed is not None:
|
||||
torch.manual_seed(self.training_seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(self.training_seed)
|
||||
random.seed(self.training_seed)
|
||||
|
||||
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 +62,18 @@ 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):
|
||||
os.makedirs(self.save_root, exist_ok=True)
|
||||
save_dif = os.path.join(self.save_root, f'config.yaml')
|
||||
with open(save_dif, 'w') as f:
|
||||
yaml.dump(self.job.raw_config, f)
|
||||
|
||||
107
jobs/process/GenerateProcess.py
Normal file
107
jobs/process/GenerateProcess.py
Normal file
@@ -0,0 +1,107 @@
|
||||
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
|
||||
import random
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.prompts: List[str]
|
||||
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")
|
||||
|
||||
if kwargs.get('shuffle', False):
|
||||
# shuffle the prompts
|
||||
random.shuffle(self.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, sampler=self.generate_config.sampler)
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
104
jobs/process/ModRescaleLoraProcess.py
Normal file
104
jobs/process/ModRescaleLoraProcess.py
Normal file
@@ -0,0 +1,104 @@
|
||||
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.process_id: int
|
||||
self.config: OrderedDict
|
||||
self.progress_bar: ForwardRef('tqdm') = None
|
||||
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}")
|
||||
657
jobs/process/TrainESRGANProcess.py
Normal file
657
jobs/process/TrainESRGANProcess.py
Normal file
@@ -0,0 +1,657 @@
|
||||
import copy
|
||||
import glob
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import List, Optional
|
||||
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
|
||||
from toolkit.basic import flush
|
||||
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.float32
|
||||
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)
|
||||
self._pattern_loss = self._pattern_loss.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, batch: Optional[List[torch.Tensor]] = None):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
os.makedirs(sample_folder, exist_ok=True)
|
||||
batch_sample_folder = os.path.join(self.save_root, 'samples_batch')
|
||||
|
||||
batch_targets = None
|
||||
batch_inputs = None
|
||||
if batch is not None and not os.path.exists(batch_sample_folder):
|
||||
os.makedirs(batch_sample_folder, exist_ok=True)
|
||||
|
||||
self.model.eval()
|
||||
|
||||
def process_and_save(img, target_img, save_path):
|
||||
img = img.to(self.device, dtype=self.esrgan_dtype)
|
||||
output = self.model(img)
|
||||
# output = (output / 2 + 0.5).clamp(0, 1)
|
||||
output = output.clamp(0, 1)
|
||||
img = img.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()
|
||||
img = img.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
|
||||
|
||||
# convert to pillow image
|
||||
output = Image.fromarray((output * 255).astype(np.uint8))
|
||||
img = Image.fromarray((img * 255).astype(np.uint8))
|
||||
|
||||
if isinstance(target_img, torch.Tensor):
|
||||
# convert to pil
|
||||
target_img = target_img.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
|
||||
target_img = Image.fromarray((target_img * 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
|
||||
)
|
||||
img = img.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_img.resize((width, height))
|
||||
output = output.resize((width, height))
|
||||
img = img.resize((width, height))
|
||||
|
||||
output_img = Image.new('RGB', (width * 3, height))
|
||||
|
||||
output_img.paste(img, (0, 0))
|
||||
output_img.paste(output, (width, 0))
|
||||
output_img.paste(target_image, (width * 2, 0))
|
||||
|
||||
output_img.save(save_path)
|
||||
|
||||
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
|
||||
|
||||
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}.jpg"
|
||||
process_and_save(img, target_image, os.path.join(sample_folder, filename))
|
||||
|
||||
if batch is not None:
|
||||
batch_targets = batch[0].detach()
|
||||
batch_inputs = batch[1].detach()
|
||||
batch_targets = torch.chunk(batch_targets, batch_targets.shape[0], dim=0)
|
||||
batch_inputs = torch.chunk(batch_inputs, batch_inputs.shape[0], dim=0)
|
||||
|
||||
for i in range(len(batch_inputs)):
|
||||
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}.jpg"
|
||||
process_and_save(batch_inputs[i], batch_targets[i], os.path.join(batch_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()
|
||||
steps_per_step = (self.critic.num_critic_per_gen + 1)
|
||||
|
||||
max_step_epochs = self.max_steps // (len(self.data_loader) // steps_per_step)
|
||||
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 * steps_per_step
|
||||
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
|
||||
critic_losses = []
|
||||
for epoch in range(self.epoch_num, self.epochs, 1):
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
flush()
|
||||
for targets, inputs in self.data_loader:
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
with torch.no_grad():
|
||||
is_critic_only_step = False
|
||||
if self.use_critic and 1 / (self.critic.num_critic_per_gen + 1) < np.random.uniform():
|
||||
is_critic_only_step = True
|
||||
|
||||
targets = targets.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1).detach()
|
||||
inputs = inputs.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1).detach()
|
||||
|
||||
optimizer.zero_grad()
|
||||
# dont do grads here for critic step
|
||||
do_grad = not is_critic_only_step
|
||||
with torch.set_grad_enabled(do_grad):
|
||||
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)
|
||||
if torch.isnan(pred).any():
|
||||
raise ValueError('pred has nan values')
|
||||
if torch.isnan(targets).any():
|
||||
raise ValueError('targets has nan values')
|
||||
|
||||
# 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)
|
||||
# make sure we dont have nans
|
||||
if torch.isnan(self.vgg19_pool_4.tensor).any():
|
||||
raise ValueError('vgg19_pool_4 has nan values')
|
||||
|
||||
if is_critic_only_step:
|
||||
critic_d_loss = self.critic.step(self.vgg19_pool_4.tensor.detach())
|
||||
critic_losses.append(critic_d_loss)
|
||||
# don't do generator step
|
||||
continue
|
||||
else:
|
||||
# doing a regular step
|
||||
if len(critic_losses) == 0:
|
||||
critic_d_loss = 0
|
||||
else:
|
||||
critic_d_loss = sum(critic_losses) / len(critic_losses)
|
||||
|
||||
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
|
||||
# make sure non nan
|
||||
if torch.isnan(loss):
|
||||
raise ValueError('loss is nan')
|
||||
|
||||
# Backward pass and optimization
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
|
||||
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, batch=[targets, inputs])
|
||||
|
||||
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()
|
||||
@@ -1,76 +0,0 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
|
||||
from toolkit.config_modules import SliderConfig
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
import sys
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
|
||||
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
|
||||
import gc
|
||||
|
||||
import torch
|
||||
from leco import train_util, model_util
|
||||
from leco.prompt_util import PromptEmbedsCache
|
||||
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
class LoRAHack:
|
||||
def __init__(self, **kwargs):
|
||||
self.type = kwargs.get('type', 'suppression')
|
||||
|
||||
|
||||
class TrainLoRAHack(BaseSDTrainProcess):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
self.hack_config = LoRAHack(**self.get_conf('hack', {}))
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
# we don't need text encoder so move it to cpu
|
||||
self.sd.text_encoder.to("cpu")
|
||||
flush()
|
||||
# end hook_before_train_loop
|
||||
|
||||
if self.hack_config.type == 'suppression':
|
||||
# set all params to self.current_suppression
|
||||
params = self.network.parameters()
|
||||
for param in params:
|
||||
# get random noise for each param
|
||||
noise = torch.randn_like(param) - 0.5
|
||||
# apply noise to param
|
||||
param.data = noise * 0.001
|
||||
|
||||
|
||||
def supress_loop(self):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'sup': 0.0}
|
||||
)
|
||||
# increase noise
|
||||
for param in self.network.parameters():
|
||||
# get random noise for each param
|
||||
noise = torch.randn_like(param) - 0.5
|
||||
# apply noise to param
|
||||
param.data = param.data + noise * 0.001
|
||||
|
||||
|
||||
|
||||
return loss_dict
|
||||
|
||||
def hook_train_loop(self):
|
||||
if self.hack_config.type == 'suppression':
|
||||
return self.supress_loop()
|
||||
else:
|
||||
raise NotImplementedError(f'unknown hack type: {self.hack_config.type}')
|
||||
# end hook_train_loop
|
||||
@@ -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,29 @@
|
||||
# ref:
|
||||
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
import copy
|
||||
import os
|
||||
from typing import Optional
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import Union
|
||||
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from torchvision.transforms import transforms
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.basic import value_map
|
||||
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
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.sd_device_states_presets import get_train_sd_device_state_preset
|
||||
from toolkit.train_tools import get_torch_dtype, apply_snr_weight, apply_learnable_snr_gos
|
||||
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 +31,15 @@ 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
|
||||
adapter_transforms = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
|
||||
|
||||
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 +48,119 @@ 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
|
||||
|
||||
# check if we have more targets than steps
|
||||
# this can happen because of permutation son shuffling
|
||||
if len(self.slider_config.targets) > self.train_config.steps:
|
||||
# trim targets
|
||||
self.slider_config.targets = self.slider_config.targets[:self.train_config.steps]
|
||||
|
||||
# get presets
|
||||
self.eval_slider_device_state = get_train_sd_device_state_preset(
|
||||
self.device_torch,
|
||||
train_unet=False,
|
||||
train_text_encoder=False,
|
||||
cached_latents=self.is_latents_cached,
|
||||
train_lora=False,
|
||||
train_adapter=False,
|
||||
train_embedding=False,
|
||||
)
|
||||
|
||||
self.train_slider_device_state = get_train_sd_device_state_preset(
|
||||
self.device_torch,
|
||||
train_unet=self.train_config.train_unet,
|
||||
train_text_encoder=False,
|
||||
cached_latents=self.is_latents_cached,
|
||||
train_lora=True,
|
||||
train_adapter=False,
|
||||
train_embedding=False,
|
||||
)
|
||||
|
||||
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] = []
|
||||
print(f"Building prompt cache")
|
||||
|
||||
# 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
|
||||
),
|
||||
]
|
||||
# trim to max steps if max steps is lower than prompt count
|
||||
# todo, this can break if we have more targets than steps, should be fixed, by reducing permuations, but could stil happen with low steps
|
||||
# prompts_to_cache = prompts_to_cache[:self.train_config.steps]
|
||||
|
||||
# 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 +173,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,78 +204,197 @@ 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()
|
||||
if self.data_loader is not None:
|
||||
# we will have images, prep the vae
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
# end hook_before_train_loop
|
||||
|
||||
def hook_train_loop(self):
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
def before_dataset_load(self):
|
||||
if self.slider_config.use_adapter == 'depth':
|
||||
print(f"Loading T2I Adapter for depth")
|
||||
# called before LoRA network is loaded but after model is loaded
|
||||
# attach the adapter here so it is there before we load the network
|
||||
adapter_path = 'TencentARC/t2iadapter_depth_sd15v2'
|
||||
if self.model_config.is_xl:
|
||||
adapter_path = 'TencentARC/t2i-adapter-depth-midas-sdxl-1.0'
|
||||
|
||||
# get a random pair
|
||||
prompt_pair: EncodedPromptPair = self.prompt_pairs[
|
||||
torch.randint(0, len(self.prompt_pairs), (1,)).item()
|
||||
]
|
||||
print(f"Loading T2I Adapter from {adapter_path}")
|
||||
|
||||
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
|
||||
# dont name this adapter since we are not training it
|
||||
self.t2i_adapter = T2IAdapter.from_pretrained(
|
||||
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype), varient="fp16"
|
||||
).to(self.device_torch)
|
||||
self.t2i_adapter.eval()
|
||||
self.t2i_adapter.requires_grad_(False)
|
||||
flush()
|
||||
|
||||
@torch.no_grad()
|
||||
def get_adapter_images(self, batch: Union[None, 'DataLoaderBatchDTO']):
|
||||
|
||||
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
adapter_folder_path = self.slider_config.adapter_img_dir
|
||||
adapter_images = []
|
||||
# loop through images
|
||||
for file_item in batch.file_items:
|
||||
img_path = file_item.path
|
||||
file_name_no_ext = os.path.basename(img_path).split('.')[0]
|
||||
# find the image
|
||||
for ext in img_ext_list:
|
||||
if os.path.exists(os.path.join(adapter_folder_path, file_name_no_ext + ext)):
|
||||
adapter_images.append(os.path.join(adapter_folder_path, file_name_no_ext + ext))
|
||||
break
|
||||
width, height = batch.file_items[0].crop_width, batch.file_items[0].crop_height
|
||||
adapter_tensors = []
|
||||
# load images with torch transforms
|
||||
for idx, adapter_image in enumerate(adapter_images):
|
||||
# we need to centrally crop the largest dimension of the image to match the batch shape after scaling
|
||||
# to the smallest dimension
|
||||
img: Image.Image = Image.open(adapter_image)
|
||||
if img.width > img.height:
|
||||
# scale down so height is the same as batch
|
||||
new_height = height
|
||||
new_width = int(img.width * (height / img.height))
|
||||
else:
|
||||
new_width = width
|
||||
new_height = int(img.height * (width / img.width))
|
||||
|
||||
img = img.resize((new_width, new_height))
|
||||
crop_fn = transforms.CenterCrop((height, width))
|
||||
# crop the center to match batch
|
||||
img = crop_fn(img)
|
||||
img = adapter_transforms(img)
|
||||
adapter_tensors.append(img)
|
||||
|
||||
# stack them
|
||||
adapter_tensors = torch.stack(adapter_tensors).to(
|
||||
self.device_torch, dtype=get_torch_dtype(self.train_config.dtype)
|
||||
)
|
||||
return adapter_tensors
|
||||
|
||||
def hook_train_loop(self, batch: Union['DataLoaderBatchDTO', None]):
|
||||
# set to eval mode
|
||||
self.sd.set_device_state(self.eval_slider_device_state)
|
||||
with torch.no_grad():
|
||||
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)
|
||||
|
||||
# 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(
|
||||
pred_kwargs = {}
|
||||
|
||||
def get_noise_pred(neg, pos, gs, cts, dn):
|
||||
down_kwargs = copy.deepcopy(pred_kwargs)
|
||||
if 'down_block_additional_residuals' in down_kwargs:
|
||||
dbr_batch_size = down_kwargs['down_block_additional_residuals'][0].shape[0]
|
||||
if dbr_batch_size != dn.shape[0]:
|
||||
amount_to_add = int(dn.shape[0] * 2 / dbr_batch_size)
|
||||
down_kwargs['down_block_additional_residuals'] = [
|
||||
torch.cat([sample.clone()] * amount_to_add) for sample in
|
||||
down_kwargs['down_block_additional_residuals']
|
||||
]
|
||||
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,
|
||||
**down_kwargs
|
||||
)
|
||||
|
||||
# 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
|
||||
)
|
||||
adapter_images = None
|
||||
self.sd.unet.eval()
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
# 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
|
||||
from_batch = False
|
||||
if batch is not None:
|
||||
# traing from a batch of images, not generating ourselves
|
||||
from_batch = True
|
||||
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
|
||||
if self.slider_config.adapter_img_dir is not None:
|
||||
adapter_images = self.get_adapter_images(batch)
|
||||
adapter_strength_min = 0.9
|
||||
adapter_strength_max = 1.0
|
||||
|
||||
# ger a random number of steps
|
||||
timesteps_to = torch.randint(
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
).item()
|
||||
def rand_strength(sample):
|
||||
adapter_conditioning_scale = torch.rand(
|
||||
(1,), device=self.device_torch, dtype=dtype
|
||||
)
|
||||
|
||||
# get noise
|
||||
noise = self.get_latent_noise(
|
||||
pixel_height=height,
|
||||
pixel_width=width,
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
adapter_conditioning_scale = value_map(
|
||||
adapter_conditioning_scale,
|
||||
0.0,
|
||||
1.0,
|
||||
adapter_strength_min,
|
||||
adapter_strength_max
|
||||
)
|
||||
return sample.to(self.device_torch, dtype=dtype).detach() * adapter_conditioning_scale
|
||||
|
||||
# get latents
|
||||
latents = noise * self.sd.noise_scheduler.init_noise_sigma
|
||||
latents = latents.to(self.device_torch, dtype=dtype)
|
||||
down_block_additional_residuals = self.t2i_adapter(adapter_images)
|
||||
down_block_additional_residuals = [
|
||||
rand_strength(sample) for sample in down_block_additional_residuals
|
||||
]
|
||||
pred_kwargs['down_block_additional_residuals'] = down_block_additional_residuals
|
||||
|
||||
with self.network:
|
||||
assert self.network.is_active
|
||||
self.network.multiplier = multiplier
|
||||
denoised_latents = self.diffuse_some_steps(
|
||||
denoised_latents = torch.cat([noisy_latents] * self.prompt_chunk_size, dim=0)
|
||||
current_timestep = timesteps
|
||||
else:
|
||||
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
|
||||
# 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=true_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)
|
||||
|
||||
assert not self.network.is_active
|
||||
self.sd.unet.eval()
|
||||
# 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,101 +402,281 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
guidance_scale=3,
|
||||
)
|
||||
|
||||
noise_scheduler.set_timesteps(1000)
|
||||
|
||||
current_timestep = noise_scheduler.timesteps[
|
||||
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
||||
]
|
||||
noise_scheduler.set_timesteps(1000)
|
||||
|
||||
current_timestep_index = int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
|
||||
current_timestep = noise_scheduler.timesteps[current_timestep_index]
|
||||
|
||||
# split the latents into out prompt pair chunks
|
||||
denoised_latent_chunks = torch.chunk(denoised_latents, self.prompt_chunk_size, dim=0)
|
||||
denoised_latent_chunks = [x.detach() for x in denoised_latent_chunks]
|
||||
|
||||
# flush() # 4.2GB to 3GB on 512x512
|
||||
mask_multiplier = torch.ones((denoised_latents.shape[0], 1, 1, 1), device=self.device_torch, dtype=dtype)
|
||||
has_mask = False
|
||||
if batch and batch.mask_tensor is not None:
|
||||
with self.timer('get_mask_multiplier'):
|
||||
# upsampling no supported for bfloat16
|
||||
mask_multiplier = batch.mask_tensor.to(self.device_torch, dtype=torch.float16).detach()
|
||||
# scale down to the size of the latents, mask multiplier shape(bs, 1, width, height), noisy_latents shape(bs, channels, width, height)
|
||||
mask_multiplier = torch.nn.functional.interpolate(
|
||||
mask_multiplier, size=(noisy_latents.shape[2], noisy_latents.shape[3])
|
||||
)
|
||||
# expand to match latents
|
||||
mask_multiplier = mask_multiplier.expand(-1, noisy_latents.shape[1], -1, -1)
|
||||
mask_multiplier = mask_multiplier.to(self.device_torch, dtype=dtype).detach()
|
||||
has_mask = True
|
||||
|
||||
if has_mask:
|
||||
unmasked_target = get_noise_pred(
|
||||
prompt_pair.positive_target, # negative prompt
|
||||
prompt_pair.target_class, # positive prompt
|
||||
1,
|
||||
current_timestep,
|
||||
denoised_latents
|
||||
)
|
||||
unmasked_target = unmasked_target.detach()
|
||||
unmasked_target.requires_grad = False
|
||||
else:
|
||||
unmasked_target = None
|
||||
|
||||
# 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 = positive_latents.detach()
|
||||
positive_latents.requires_grad = False
|
||||
|
||||
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 = neutral_latents.detach()
|
||||
neutral_latents.requires_grad = False
|
||||
|
||||
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 = unconditional_latents.detach()
|
||||
unconditional_latents.requires_grad = False
|
||||
|
||||
offset = guidance_scale * (positive_latents - unconditional_latents)
|
||||
denoised_latents = denoised_latents.detach()
|
||||
|
||||
offset_neutral = neutral_latents
|
||||
if erase:
|
||||
offset_neutral -= offset
|
||||
else:
|
||||
# enhance
|
||||
offset_neutral += offset
|
||||
self.sd.set_device_state(self.train_slider_device_state)
|
||||
self.sd.unet.train()
|
||||
# start accumulating gradients
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
loss = loss_function(
|
||||
target_latents,
|
||||
offset_neutral,
|
||||
) * weight
|
||||
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_slide = loss.item()
|
||||
# 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)
|
||||
|
||||
if anchor_loss is not None:
|
||||
loss += anchor_loss
|
||||
# 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)
|
||||
|
||||
loss_float = loss.item()
|
||||
# 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 = loss.to(self.device_torch)
|
||||
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()
|
||||
|
||||
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")
|
||||
|
||||
with torch.no_grad():
|
||||
if self.slider_config.low_ram:
|
||||
prompt_pair_chunks = split_prompt_pairs(prompt_pair.detach(), self.prompt_chunk_size)
|
||||
denoised_latent_chunks = denoised_latent_chunks # just to have it in one place
|
||||
positive_latents_chunks = torch.chunk(positive_latents.detach(), self.prompt_chunk_size, dim=0)
|
||||
neutral_latents_chunks = torch.chunk(neutral_latents.detach(), self.prompt_chunk_size, dim=0)
|
||||
unconditional_latents_chunks = torch.chunk(
|
||||
unconditional_latents.detach(),
|
||||
self.prompt_chunk_size,
|
||||
dim=0
|
||||
)
|
||||
mask_multiplier_chunks = torch.chunk(mask_multiplier, self.prompt_chunk_size, dim=0)
|
||||
if unmasked_target is not None:
|
||||
unmasked_target_chunks = torch.chunk(unmasked_target, self.prompt_chunk_size, dim=0)
|
||||
else:
|
||||
unmasked_target_chunks = [None for _ in range(self.prompt_chunk_size)]
|
||||
else:
|
||||
# run through in one instance
|
||||
prompt_pair_chunks = [prompt_pair.detach()]
|
||||
denoised_latent_chunks = [torch.cat(denoised_latent_chunks, dim=0).detach()]
|
||||
positive_latents_chunks = [positive_latents.detach()]
|
||||
neutral_latents_chunks = [neutral_latents.detach()]
|
||||
unconditional_latents_chunks = [unconditional_latents.detach()]
|
||||
mask_multiplier_chunks = [mask_multiplier]
|
||||
unmasked_target_chunks = [unmasked_target]
|
||||
|
||||
# flush()
|
||||
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, \
|
||||
mask_multiplier_chunk, \
|
||||
unmasked_target_chunk \
|
||||
in zip(
|
||||
prompt_pair_chunks,
|
||||
denoised_latent_chunks,
|
||||
positive_latents_chunks,
|
||||
neutral_latents_chunks,
|
||||
unconditional_latents_chunks,
|
||||
mask_multiplier_chunks,
|
||||
unmasked_target_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 = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")
|
||||
|
||||
# do inverted mask to preserve non masked
|
||||
if has_mask and unmasked_target_chunk is not None:
|
||||
loss = loss * mask_multiplier_chunk
|
||||
# match the mask unmasked_target_chunk
|
||||
mask_target_loss = torch.nn.functional.mse_loss(
|
||||
target_latents.float(),
|
||||
unmasked_target_chunk.float(),
|
||||
reduction="none"
|
||||
)
|
||||
mask_target_loss = mask_target_loss * (1.0 - mask_multiplier_chunk)
|
||||
loss += mask_target_loss
|
||||
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
if self.train_config.learnable_snr_gos:
|
||||
if from_batch:
|
||||
# match batch size
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler,
|
||||
self.train_config.min_snr_gamma)
|
||||
else:
|
||||
# match batch size
|
||||
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
|
||||
# add snr_gamma
|
||||
loss = apply_learnable_snr_gos(loss, timesteps_index_list, self.snr_gos)
|
||||
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
|
||||
if from_batch:
|
||||
# match batch size
|
||||
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler,
|
||||
self.train_config.min_snr_gamma)
|
||||
else:
|
||||
# match batch size
|
||||
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
|
||||
# add min_snr_gamma
|
||||
loss = apply_snr_weight(loss, timesteps_index_list, noise_scheduler,
|
||||
self.train_config.min_snr_gamma)
|
||||
|
||||
|
||||
loss = loss.mean() * 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
|
||||
)
|
||||
flush()
|
||||
# move back to cpu
|
||||
prompt_pair.to("cpu")
|
||||
# flush()
|
||||
|
||||
# reset network
|
||||
self.network.multiplier = 1.0
|
||||
@@ -399,9 +684,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,10 @@ from .BaseTrainProcess import BaseTrainProcess
|
||||
from .TrainVAEProcess import TrainVAEProcess
|
||||
from .BaseMergeProcess import BaseMergeProcess
|
||||
from .TrainSliderProcess import TrainSliderProcess
|
||||
from .TrainLoRAHack import TrainLoRAHack
|
||||
from .TrainSDRescaleProcess import TrainSDRescaleProcess
|
||||
from .TrainSliderProcessOld import TrainSliderProcessOld
|
||||
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)
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
critic_losses = []
|
||||
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).float()
|
||||
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()
|
||||
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
|
||||
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
|
||||
|
||||
|
||||
339
notebooks/SliderTraining.ipynb
Normal file
339
notebooks/SliderTraining.ipynb
Normal file
@@ -0,0 +1,339 @@
|
||||
{
|
||||
"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": "code",
|
||||
"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
|
||||
},
|
||||
"outputs": []
|
||||
},
|
||||
{
|
||||
"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"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
1
repositories/batch_annotator
Submodule
1
repositories/batch_annotator
Submodule
Submodule repositories/batch_annotator added at 420e142f6a
1
repositories/ipadapter
Submodule
1
repositories/ipadapter
Submodule
Submodule repositories/ipadapter added at d8ab37c421
@@ -1,11 +1,10 @@
|
||||
torch
|
||||
torchvision
|
||||
safetensors
|
||||
diffusers
|
||||
transformers
|
||||
lycoris_lora
|
||||
diffusers==0.21.3
|
||||
git+https://github.com/huggingface/transformers.git
|
||||
lycoris-lora==1.8.3
|
||||
flatten_json
|
||||
accelerator
|
||||
pyyaml
|
||||
oyaml
|
||||
tensorboard
|
||||
@@ -14,4 +13,12 @@ invisible-watermark
|
||||
einops
|
||||
accelerate
|
||||
toml
|
||||
albumentations
|
||||
albumentations
|
||||
pydantic
|
||||
omegaconf
|
||||
k-diffusion
|
||||
open_clip_torch
|
||||
timm
|
||||
prodigyopt
|
||||
controlnet_aux==0.0.7
|
||||
python-dotenv
|
||||
26
run.py
26
run.py
@@ -1,6 +1,22 @@
|
||||
import os
|
||||
import sys
|
||||
from typing import Union, OrderedDict
|
||||
from dotenv import load_dotenv
|
||||
# Load the .env file if it exists
|
||||
load_dotenv()
|
||||
|
||||
sys.path.insert(0, os.getcwd())
|
||||
# must come before ANY torch or fastai imports
|
||||
# import toolkit.cuda_malloc
|
||||
|
||||
# turn off diffusers telemetry until I can figure out how to make it opt-in
|
||||
os.environ['DISABLE_TELEMETRY'] = 'YES'
|
||||
|
||||
# check if we have DEBUG_TOOLKIT in env
|
||||
if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
# set torch to trace mode
|
||||
import torch
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
|
||||
@@ -36,6 +52,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 +73,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
|
||||
|
||||
128
scripts/convert_cog.py
Normal file
128
scripts/convert_cog.py
Normal file
@@ -0,0 +1,128 @@
|
||||
import json
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import save_file
|
||||
|
||||
device = torch.device('cpu')
|
||||
|
||||
# [diffusers] -> kohya
|
||||
embedding_mapping = {
|
||||
'text_encoders_0': 'clip_l',
|
||||
'text_encoders_1': 'clip_g'
|
||||
}
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
KEYMAP_ROOT = os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps')
|
||||
sdxl_keymap_path = os.path.join(KEYMAP_ROOT, 'stable_diffusion_locon_sdxl.json')
|
||||
|
||||
# load keymap
|
||||
with open(sdxl_keymap_path, 'r') as f:
|
||||
ldm_diffusers_keymap = json.load(f)['ldm_diffusers_keymap']
|
||||
|
||||
# invert the item / key pairs
|
||||
diffusers_ldm_keymap = {v: k for k, v in ldm_diffusers_keymap.items()}
|
||||
|
||||
|
||||
def get_ldm_key(diffuser_key):
|
||||
diffuser_key = f"lora_unet_{diffuser_key.replace('.', '_')}"
|
||||
diffuser_key = diffuser_key.replace('_lora_down_weight', '.lora_down.weight')
|
||||
diffuser_key = diffuser_key.replace('_lora_up_weight', '.lora_up.weight')
|
||||
diffuser_key = diffuser_key.replace('_alpha', '.alpha')
|
||||
diffuser_key = diffuser_key.replace('_processor_to_', '_to_')
|
||||
diffuser_key = diffuser_key.replace('_to_out.', '_to_out_0.')
|
||||
if diffuser_key in diffusers_ldm_keymap:
|
||||
return diffusers_ldm_keymap[diffuser_key]
|
||||
else:
|
||||
raise KeyError(f"Key {diffuser_key} not found in keymap")
|
||||
|
||||
|
||||
def convert_cog(lora_path, embedding_path):
|
||||
embedding_state_dict = OrderedDict()
|
||||
lora_state_dict = OrderedDict()
|
||||
|
||||
# # normal dict
|
||||
# normal_dict = OrderedDict()
|
||||
# example_path = "/mnt/Models/stable-diffusion/models/LoRA/sdxl/LogoRedmond_LogoRedAF.safetensors"
|
||||
# with safe_open(example_path, framework="pt", device='cpu') as f:
|
||||
# keys = list(f.keys())
|
||||
# for key in keys:
|
||||
# normal_dict[key] = f.get_tensor(key)
|
||||
|
||||
with safe_open(embedding_path, framework="pt", device='cpu') as f:
|
||||
keys = list(f.keys())
|
||||
for key in keys:
|
||||
new_key = embedding_mapping[key]
|
||||
embedding_state_dict[new_key] = f.get_tensor(key)
|
||||
|
||||
with safe_open(lora_path, framework="pt", device='cpu') as f:
|
||||
keys = list(f.keys())
|
||||
lora_rank = None
|
||||
|
||||
# get the lora dim first. Check first 3 linear layers just to be safe
|
||||
for key in keys:
|
||||
new_key = get_ldm_key(key)
|
||||
tensor = f.get_tensor(key)
|
||||
num_checked = 0
|
||||
if len(tensor.shape) == 2:
|
||||
this_dim = min(tensor.shape)
|
||||
if lora_rank is None:
|
||||
lora_rank = this_dim
|
||||
elif lora_rank != this_dim:
|
||||
raise ValueError(f"lora rank is not consistent, got {tensor.shape}")
|
||||
else:
|
||||
num_checked += 1
|
||||
if num_checked >= 3:
|
||||
break
|
||||
|
||||
for key in keys:
|
||||
new_key = get_ldm_key(key)
|
||||
tensor = f.get_tensor(key)
|
||||
if new_key.endswith('.lora_down.weight'):
|
||||
alpha_key = new_key.replace('.lora_down.weight', '.alpha')
|
||||
# diffusers does not have alpha, they usa an alpha multiplier of 1 which is a tensor weight of the dims
|
||||
# assume first smallest dim is the lora rank if shape is 2
|
||||
lora_state_dict[alpha_key] = torch.ones(1).to(tensor.device, tensor.dtype) * lora_rank
|
||||
|
||||
lora_state_dict[new_key] = tensor
|
||||
|
||||
return lora_state_dict, embedding_state_dict
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
'lora_path',
|
||||
type=str,
|
||||
help='Path to lora file'
|
||||
)
|
||||
parser.add_argument(
|
||||
'embedding_path',
|
||||
type=str,
|
||||
help='Path to embedding file'
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'--lora_output',
|
||||
type=str,
|
||||
default="lora_output",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'--embedding_output',
|
||||
type=str,
|
||||
default="embedding_output",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
lora_state_dict, embedding_state_dict = convert_cog(args.lora_path, args.embedding_path)
|
||||
|
||||
# save them
|
||||
save_file(lora_state_dict, args.lora_output)
|
||||
save_file(embedding_state_dict, args.embedding_output)
|
||||
print(f"Saved lora to {args.lora_output}")
|
||||
print(f"Saved embedding to {args.embedding_output}")
|
||||
57
scripts/make_diffusers_model.py
Normal file
57
scripts/make_diffusers_model.py
Normal file
@@ -0,0 +1,57 @@
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
'input_path',
|
||||
type=str,
|
||||
help='Path to original sdxl model'
|
||||
)
|
||||
parser.add_argument(
|
||||
'output_path',
|
||||
type=str,
|
||||
help='output path'
|
||||
)
|
||||
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
|
||||
parser.add_argument('--refiner', action='store_true', help='is refiner model')
|
||||
parser.add_argument('--ssd', action='store_true', help='is ssd model')
|
||||
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
|
||||
|
||||
args = parser.parse_args()
|
||||
device = torch.device('cpu')
|
||||
dtype = torch.float32
|
||||
|
||||
print(f"Loading model from {args.input_path}")
|
||||
|
||||
|
||||
diffusers_model_config = ModelConfig(
|
||||
name_or_path=args.input_path,
|
||||
is_xl=args.sdxl,
|
||||
is_v2=args.sd2,
|
||||
is_ssd=args.ssd,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd = StableDiffusion(
|
||||
model_config=diffusers_model_config,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd.load_model()
|
||||
|
||||
|
||||
print(f"Loaded model from {args.input_path}")
|
||||
|
||||
diffusers_sd.pipeline.fuse_lora()
|
||||
|
||||
meta = OrderedDict()
|
||||
|
||||
diffusers_sd.save(args.output_path, meta=meta)
|
||||
|
||||
|
||||
print(f"Saved to {args.output_path}")
|
||||
67
scripts/make_lcm_sdxl_model.py
Normal file
67
scripts/make_lcm_sdxl_model.py
Normal file
@@ -0,0 +1,67 @@
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
'input_path',
|
||||
type=str,
|
||||
help='Path to original sdxl model'
|
||||
)
|
||||
parser.add_argument(
|
||||
'output_path',
|
||||
type=str,
|
||||
help='output path'
|
||||
)
|
||||
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
|
||||
parser.add_argument('--refiner', action='store_true', help='is refiner model')
|
||||
parser.add_argument('--ssd', action='store_true', help='is ssd model')
|
||||
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
|
||||
|
||||
args = parser.parse_args()
|
||||
device = torch.device('cpu')
|
||||
dtype = torch.float32
|
||||
|
||||
print(f"Loading model from {args.input_path}")
|
||||
|
||||
if args.sdxl:
|
||||
adapter_id = "latent-consistency/lcm-lora-sdxl"
|
||||
if args.refiner:
|
||||
adapter_id = "latent-consistency/lcm-lora-sdxl"
|
||||
elif args.ssd:
|
||||
adapter_id = "latent-consistency/lcm-lora-ssd-1b"
|
||||
else:
|
||||
adapter_id = "latent-consistency/lcm-lora-sdv1-5"
|
||||
|
||||
|
||||
diffusers_model_config = ModelConfig(
|
||||
name_or_path=args.input_path,
|
||||
is_xl=args.sdxl,
|
||||
is_v2=args.sd2,
|
||||
is_ssd=args.ssd,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd = StableDiffusion(
|
||||
model_config=diffusers_model_config,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd.load_model()
|
||||
|
||||
|
||||
print(f"Loaded model from {args.input_path}")
|
||||
|
||||
diffusers_sd.pipeline.load_lora_weights(adapter_id)
|
||||
diffusers_sd.pipeline.fuse_lora()
|
||||
|
||||
meta = OrderedDict()
|
||||
|
||||
diffusers_sd.save(args.output_path, meta=meta)
|
||||
|
||||
|
||||
print(f"Saved to {args.output_path}")
|
||||
@@ -1,547 +0,0 @@
|
||||
import gc
|
||||
import time
|
||||
import argparse
|
||||
import itertools
|
||||
import math
|
||||
import os
|
||||
from multiprocessing import Value
|
||||
|
||||
from tqdm import tqdm
|
||||
import torch
|
||||
from accelerate.utils import set_seed
|
||||
import diffusers
|
||||
from diffusers import DDPMScheduler
|
||||
|
||||
import library.train_util as train_util
|
||||
import library.config_util as config_util
|
||||
from library.config_util import (
|
||||
ConfigSanitizer,
|
||||
BlueprintGenerator,
|
||||
)
|
||||
import custom_tools.train_tools as train_tools
|
||||
import library.custom_train_functions as custom_train_functions
|
||||
from library.custom_train_functions import (
|
||||
apply_snr_weight,
|
||||
get_weighted_text_embeddings,
|
||||
prepare_scheduler_for_custom_training,
|
||||
pyramid_noise_like,
|
||||
apply_noise_offset,
|
||||
scale_v_prediction_loss_like_noise_prediction,
|
||||
)
|
||||
|
||||
# perlin_noise,
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
SD_SCRIPTS_ROOT = os.path.join(PROJECT_ROOT, "repositories", "sd-scripts")
|
||||
|
||||
|
||||
def train(args):
|
||||
train_util.verify_training_args(args)
|
||||
train_util.prepare_dataset_args(args, False)
|
||||
|
||||
cache_latents = args.cache_latents
|
||||
|
||||
if args.seed is not None:
|
||||
set_seed(args.seed) # 乱数系列を初期化する
|
||||
|
||||
tokenizer = train_util.load_tokenizer(args)
|
||||
|
||||
# データセットを準備する
|
||||
if args.dataset_class is None:
|
||||
blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, False, True))
|
||||
if args.dataset_config is not None:
|
||||
print(f"Load dataset config from {args.dataset_config}")
|
||||
user_config = config_util.load_user_config(args.dataset_config)
|
||||
ignored = ["train_data_dir", "reg_data_dir"]
|
||||
if any(getattr(args, attr) is not None for attr in ignored):
|
||||
print(
|
||||
"ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format(
|
||||
", ".join(ignored)
|
||||
)
|
||||
)
|
||||
else:
|
||||
user_config = {
|
||||
"datasets": [
|
||||
{"subsets": config_util.generate_dreambooth_subsets_config_by_subdirs(args.train_data_dir, args.reg_data_dir)}
|
||||
]
|
||||
}
|
||||
|
||||
blueprint = blueprint_generator.generate(user_config, args, tokenizer=tokenizer)
|
||||
train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group)
|
||||
else:
|
||||
train_dataset_group = train_util.load_arbitrary_dataset(args, tokenizer)
|
||||
|
||||
current_epoch = Value("i", 0)
|
||||
current_step = Value("i", 0)
|
||||
ds_for_collater = train_dataset_group if args.max_data_loader_n_workers == 0 else None
|
||||
collater = train_util.collater_class(current_epoch, current_step, ds_for_collater)
|
||||
|
||||
if args.no_token_padding:
|
||||
train_dataset_group.disable_token_padding()
|
||||
|
||||
if args.debug_dataset:
|
||||
train_util.debug_dataset(train_dataset_group)
|
||||
return
|
||||
|
||||
if cache_latents:
|
||||
assert (
|
||||
train_dataset_group.is_latent_cacheable()
|
||||
), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません"
|
||||
|
||||
# replace captions with names
|
||||
if args.name_replace is not None:
|
||||
print(f"Replacing captions [name] with '{args.name_replace}'")
|
||||
|
||||
train_dataset_group = train_tools.replace_filewords_in_dataset_group(
|
||||
train_dataset_group, args
|
||||
)
|
||||
|
||||
# acceleratorを準備する
|
||||
print("prepare accelerator")
|
||||
|
||||
if args.gradient_accumulation_steps > 1:
|
||||
print(
|
||||
f"gradient_accumulation_steps is {args.gradient_accumulation_steps}. accelerate does not support gradient_accumulation_steps when training multiple models (U-Net and Text Encoder), so something might be wrong"
|
||||
)
|
||||
print(
|
||||
f"gradient_accumulation_stepsが{args.gradient_accumulation_steps}に設定されています。accelerateは複数モデル(U-NetおよびText Encoder)の学習時にgradient_accumulation_stepsをサポートしていないため結果は未知数です"
|
||||
)
|
||||
|
||||
accelerator, unwrap_model = train_util.prepare_accelerator(args)
|
||||
|
||||
# mixed precisionに対応した型を用意しておき適宜castする
|
||||
weight_dtype, save_dtype = train_util.prepare_dtype(args)
|
||||
|
||||
# モデルを読み込む
|
||||
text_encoder, vae, unet, load_stable_diffusion_format = train_util.load_target_model(args, weight_dtype, accelerator)
|
||||
|
||||
# verify load/save model formats
|
||||
if load_stable_diffusion_format:
|
||||
src_stable_diffusion_ckpt = args.pretrained_model_name_or_path
|
||||
src_diffusers_model_path = None
|
||||
else:
|
||||
src_stable_diffusion_ckpt = None
|
||||
src_diffusers_model_path = args.pretrained_model_name_or_path
|
||||
|
||||
if args.save_model_as is None:
|
||||
save_stable_diffusion_format = load_stable_diffusion_format
|
||||
use_safetensors = args.use_safetensors
|
||||
else:
|
||||
save_stable_diffusion_format = args.save_model_as.lower() == "ckpt" or args.save_model_as.lower() == "safetensors"
|
||||
use_safetensors = args.use_safetensors or ("safetensors" in args.save_model_as.lower())
|
||||
|
||||
# モデルに xformers とか memory efficient attention を組み込む
|
||||
train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers)
|
||||
|
||||
# 学習を準備する
|
||||
if cache_latents:
|
||||
vae.to(accelerator.device, dtype=weight_dtype)
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
with torch.no_grad():
|
||||
train_dataset_group.cache_latents(vae, args.vae_batch_size, args.cache_latents_to_disk, accelerator.is_main_process)
|
||||
vae.to("cpu")
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
# 学習を準備する:モデルを適切な状態にする
|
||||
train_text_encoder = args.stop_text_encoder_training is None or args.stop_text_encoder_training >= 0
|
||||
unet.requires_grad_(True) # 念のため追加
|
||||
text_encoder.requires_grad_(train_text_encoder)
|
||||
if not train_text_encoder:
|
||||
print("Text Encoder is not trained.")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
unet.enable_gradient_checkpointing()
|
||||
text_encoder.gradient_checkpointing_enable()
|
||||
|
||||
if not cache_latents:
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
vae.to(accelerator.device, dtype=weight_dtype)
|
||||
|
||||
# 学習に必要なクラスを準備する
|
||||
print("prepare optimizer, data loader etc.")
|
||||
if train_text_encoder:
|
||||
trainable_params = itertools.chain(unet.parameters(), text_encoder.parameters())
|
||||
else:
|
||||
trainable_params = unet.parameters()
|
||||
|
||||
_, _, optimizer = train_util.get_optimizer(args, trainable_params)
|
||||
|
||||
# dataloaderを準備する
|
||||
# DataLoaderのプロセス数:0はメインプロセスになる
|
||||
n_workers = min(args.max_data_loader_n_workers, os.cpu_count() - 1) # cpu_count-1 ただし最大で指定された数まで
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset_group,
|
||||
batch_size=1,
|
||||
shuffle=True,
|
||||
collate_fn=collater,
|
||||
num_workers=n_workers,
|
||||
persistent_workers=args.persistent_data_loader_workers,
|
||||
)
|
||||
|
||||
# 学習ステップ数を計算する
|
||||
if args.max_train_epochs is not None:
|
||||
args.max_train_steps = args.max_train_epochs * math.ceil(
|
||||
len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps
|
||||
)
|
||||
print(f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}")
|
||||
|
||||
# データセット側にも学習ステップを送信
|
||||
train_dataset_group.set_max_train_steps(args.max_train_steps)
|
||||
|
||||
if args.stop_text_encoder_training is None:
|
||||
args.stop_text_encoder_training = args.max_train_steps + 1 # do not stop until end
|
||||
|
||||
# lr schedulerを用意する TODO gradient_accumulation_stepsの扱いが何かおかしいかもしれない。後で確認する
|
||||
lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)
|
||||
|
||||
# 実験的機能:勾配も含めたfp16学習を行う モデル全体をfp16にする
|
||||
if args.full_fp16:
|
||||
assert (
|
||||
args.mixed_precision == "fp16"
|
||||
), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"
|
||||
print("enable full fp16 training.")
|
||||
unet.to(weight_dtype)
|
||||
text_encoder.to(weight_dtype)
|
||||
|
||||
# acceleratorがなんかよろしくやってくれるらしい
|
||||
if train_text_encoder:
|
||||
unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
|
||||
unet, text_encoder, optimizer, train_dataloader, lr_scheduler
|
||||
)
|
||||
else:
|
||||
unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(unet, optimizer, train_dataloader, lr_scheduler)
|
||||
|
||||
# transform DDP after prepare
|
||||
text_encoder, unet = train_util.transform_if_model_is_DDP(text_encoder, unet)
|
||||
|
||||
if not train_text_encoder:
|
||||
text_encoder.to(accelerator.device, dtype=weight_dtype) # to avoid 'cpu' vs 'cuda' error
|
||||
|
||||
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
|
||||
if args.full_fp16:
|
||||
train_util.patch_accelerator_for_fp16_training(accelerator)
|
||||
|
||||
# resumeする
|
||||
train_util.resume_from_local_or_hf_if_specified(accelerator, args)
|
||||
|
||||
# epoch数を計算する
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
||||
num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0):
|
||||
args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1
|
||||
|
||||
# 学習する
|
||||
total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
|
||||
print("running training / 学習開始")
|
||||
print(f" num train images * repeats / 学習画像の数×繰り返し回数: {train_dataset_group.num_train_images}")
|
||||
print(f" num reg images / 正則化画像の数: {train_dataset_group.num_reg_images}")
|
||||
print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}")
|
||||
print(f" num epochs / epoch数: {num_train_epochs}")
|
||||
print(f" batch size per device / バッチサイズ: {args.train_batch_size}")
|
||||
print(f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}")
|
||||
print(f" gradient ccumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}")
|
||||
print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}")
|
||||
|
||||
progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps")
|
||||
global_step = 0
|
||||
|
||||
noise_scheduler = DDPMScheduler(
|
||||
beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False
|
||||
)
|
||||
prepare_scheduler_for_custom_training(noise_scheduler, accelerator.device)
|
||||
|
||||
if accelerator.is_main_process:
|
||||
accelerator.init_trackers("dreambooth" if args.log_tracker_name is None else args.log_tracker_name)
|
||||
|
||||
if args.sample_first or args.sample_only:
|
||||
# Do initial sample before starting training
|
||||
train_tools.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizer,
|
||||
text_encoder, unet, force_sample=True)
|
||||
|
||||
if args.sample_only:
|
||||
return
|
||||
loss_list = []
|
||||
loss_total = 0.0
|
||||
for epoch in range(num_train_epochs):
|
||||
print(f"\nepoch {epoch+1}/{num_train_epochs}")
|
||||
current_epoch.value = epoch + 1
|
||||
|
||||
# 指定したステップ数までText Encoderを学習する:epoch最初の状態
|
||||
unet.train()
|
||||
# train==True is required to enable gradient_checkpointing
|
||||
if args.gradient_checkpointing or global_step < args.stop_text_encoder_training:
|
||||
text_encoder.train()
|
||||
|
||||
for step, batch in enumerate(train_dataloader):
|
||||
current_step.value = global_step
|
||||
# 指定したステップ数でText Encoderの学習を止める
|
||||
if global_step == args.stop_text_encoder_training:
|
||||
print(f"stop text encoder training at step {global_step}")
|
||||
if not args.gradient_checkpointing:
|
||||
text_encoder.train(False)
|
||||
text_encoder.requires_grad_(False)
|
||||
|
||||
with accelerator.accumulate(unet):
|
||||
with torch.no_grad():
|
||||
# latentに変換
|
||||
if cache_latents:
|
||||
latents = batch["latents"].to(accelerator.device)
|
||||
else:
|
||||
latents = vae.encode(batch["images"].to(dtype=weight_dtype)).latent_dist.sample()
|
||||
latents = latents * 0.18215
|
||||
b_size = latents.shape[0]
|
||||
|
||||
# Sample noise that we'll add to the latents
|
||||
if args.train_noise_seed is not None:
|
||||
torch.manual_seed(args.train_noise_seed)
|
||||
torch.cuda.manual_seed(args.train_noise_seed)
|
||||
# make same seed for each item in the batch by stacking them
|
||||
single_noise = torch.randn_like(latents[0])
|
||||
noise = torch.stack([single_noise for _ in range(b_size)])
|
||||
noise = noise.to(latents.device)
|
||||
elif args.seed_lock:
|
||||
noise = train_tools.get_noise_from_latents(latents)
|
||||
else:
|
||||
noise = torch.randn_like(latents, device=latents.device)
|
||||
|
||||
if args.noise_offset:
|
||||
noise = apply_noise_offset(latents, noise, args.noise_offset, args.adaptive_noise_scale)
|
||||
elif args.multires_noise_iterations:
|
||||
noise = pyramid_noise_like(noise, latents.device, args.multires_noise_iterations, args.multires_noise_discount)
|
||||
# elif args.perlin_noise:
|
||||
# noise = perlin_noise(noise, latents.device, args.perlin_noise) # only shape of noise is used currently
|
||||
|
||||
# Get the text embedding for conditioning
|
||||
with torch.set_grad_enabled(global_step < args.stop_text_encoder_training):
|
||||
if args.weighted_captions:
|
||||
encoder_hidden_states = get_weighted_text_embeddings(
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
batch["captions"],
|
||||
accelerator.device,
|
||||
args.max_token_length // 75 if args.max_token_length else 1,
|
||||
clip_skip=args.clip_skip,
|
||||
)
|
||||
else:
|
||||
input_ids = batch["input_ids"].to(accelerator.device)
|
||||
encoder_hidden_states = train_util.get_hidden_states(
|
||||
args, input_ids, tokenizer, text_encoder, None if not args.full_fp16 else weight_dtype
|
||||
)
|
||||
|
||||
# Sample a random timestep for each image
|
||||
timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (b_size,), device=latents.device)
|
||||
timesteps = timesteps.long()
|
||||
|
||||
# Add noise to the latents according to the noise magnitude at each timestep
|
||||
# (this is the forward diffusion process)
|
||||
noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
|
||||
|
||||
# Predict the noise residual
|
||||
with accelerator.autocast():
|
||||
noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample
|
||||
|
||||
if args.v_parameterization:
|
||||
# v-parameterization training
|
||||
target = noise_scheduler.get_velocity(latents, noise, timesteps)
|
||||
else:
|
||||
target = noise
|
||||
|
||||
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
|
||||
loss = loss.mean([1, 2, 3])
|
||||
|
||||
loss_weights = batch["loss_weights"] # 各sampleごとのweight
|
||||
loss = loss * loss_weights
|
||||
|
||||
if args.min_snr_gamma:
|
||||
loss = apply_snr_weight(loss, timesteps, noise_scheduler, args.min_snr_gamma)
|
||||
if args.scale_v_pred_loss_like_noise_pred:
|
||||
loss = scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler)
|
||||
|
||||
loss = loss.mean() # 平均なのでbatch_sizeで割る必要なし
|
||||
|
||||
accelerator.backward(loss)
|
||||
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
|
||||
if train_text_encoder:
|
||||
params_to_clip = itertools.chain(unet.parameters(), text_encoder.parameters())
|
||||
else:
|
||||
params_to_clip = unet.parameters()
|
||||
accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
|
||||
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
# Checks if the accelerator has performed an optimization step behind the scenes
|
||||
if accelerator.sync_gradients:
|
||||
progress_bar.update(1)
|
||||
global_step += 1
|
||||
|
||||
train_util.sample_images(
|
||||
accelerator, args, None, global_step, accelerator.device, vae, tokenizer, text_encoder, unet
|
||||
)
|
||||
|
||||
# 指定ステップごとにモデルを保存
|
||||
if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0:
|
||||
accelerator.wait_for_everyone()
|
||||
if accelerator.is_main_process:
|
||||
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||
args,
|
||||
False,
|
||||
accelerator,
|
||||
src_path,
|
||||
save_stable_diffusion_format,
|
||||
use_safetensors,
|
||||
save_dtype,
|
||||
epoch,
|
||||
num_train_epochs,
|
||||
global_step,
|
||||
unwrap_model(text_encoder),
|
||||
unwrap_model(unet),
|
||||
vae,
|
||||
)
|
||||
|
||||
current_loss = loss.detach().item()
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss": current_loss, "lr": float(lr_scheduler.get_last_lr()[0])}
|
||||
if args.optimizer_type.lower().startswith("DAdapt".lower()) or args.optimizer_type.lower() == "Prodigy".lower(): # tracking d*lr value
|
||||
logs["lr/d*lr"] = (
|
||||
lr_scheduler.optimizers[0].param_groups[0]["d"] * lr_scheduler.optimizers[0].param_groups[0]["lr"]
|
||||
)
|
||||
accelerator.log(logs, step=global_step)
|
||||
|
||||
if epoch == 0:
|
||||
loss_list.append(current_loss)
|
||||
else:
|
||||
loss_total -= loss_list[step]
|
||||
loss_list[step] = current_loss
|
||||
loss_total += current_loss
|
||||
avr_loss = loss_total / len(loss_list)
|
||||
logs = {"loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}
|
||||
progress_bar.set_postfix(**logs)
|
||||
|
||||
if global_step >= args.max_train_steps:
|
||||
break
|
||||
|
||||
if args.logging_dir is not None:
|
||||
logs = {"loss/epoch": loss_total / len(loss_list)}
|
||||
accelerator.log(logs, step=epoch + 1)
|
||||
|
||||
accelerator.wait_for_everyone()
|
||||
|
||||
if args.save_every_n_epochs is not None:
|
||||
if accelerator.is_main_process:
|
||||
# checking for saving is in util
|
||||
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||
train_util.save_sd_model_on_epoch_end_or_stepwise(
|
||||
args,
|
||||
True,
|
||||
accelerator,
|
||||
src_path,
|
||||
save_stable_diffusion_format,
|
||||
use_safetensors,
|
||||
save_dtype,
|
||||
epoch,
|
||||
num_train_epochs,
|
||||
global_step,
|
||||
unwrap_model(text_encoder),
|
||||
unwrap_model(unet),
|
||||
vae,
|
||||
)
|
||||
|
||||
train_util.sample_images(accelerator, args, epoch + 1, global_step, accelerator.device, vae, tokenizer, text_encoder, unet)
|
||||
|
||||
is_main_process = accelerator.is_main_process
|
||||
if is_main_process:
|
||||
unet = unwrap_model(unet)
|
||||
text_encoder = unwrap_model(text_encoder)
|
||||
|
||||
accelerator.end_training()
|
||||
|
||||
if args.save_state and is_main_process:
|
||||
train_util.save_state_on_train_end(args, accelerator)
|
||||
|
||||
del accelerator # この後メモリを使うのでこれは消す
|
||||
|
||||
if is_main_process:
|
||||
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
|
||||
train_util.save_sd_model_on_train_end(
|
||||
args, src_path, save_stable_diffusion_format, use_safetensors, save_dtype, epoch, global_step, text_encoder, unet, vae
|
||||
)
|
||||
print("model saved.")
|
||||
|
||||
|
||||
def setup_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
train_util.add_sd_models_arguments(parser)
|
||||
train_util.add_dataset_arguments(parser, True, False, True)
|
||||
train_util.add_training_arguments(parser, True)
|
||||
train_util.add_sd_saving_arguments(parser)
|
||||
train_util.add_optimizer_arguments(parser)
|
||||
config_util.add_config_arguments(parser)
|
||||
custom_train_functions.add_custom_train_arguments(parser)
|
||||
|
||||
parser.add_argument(
|
||||
"--no_token_padding",
|
||||
action="store_true",
|
||||
help="disable token padding (same as Diffuser's DreamBooth) / トークンのpaddingを無効にする(Diffusers版DreamBoothと同じ動作)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stop_text_encoder_training",
|
||||
type=int,
|
||||
default=None,
|
||||
help="steps to stop text encoder training, -1 for no training / Text Encoderの学習を止めるステップ数、-1で最初から学習しない",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--sample_first",
|
||||
action="store_true",
|
||||
help="Sample first interval before training",
|
||||
default=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--name_replace",
|
||||
type=str,
|
||||
help="Replaces [name] in prompts. Used is sampling, training, and regs",
|
||||
default=None
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--train_noise_seed",
|
||||
type=int,
|
||||
help="Use custom seed for training noise",
|
||||
default=None
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--sample_only",
|
||||
action="store_true",
|
||||
help="Only generate samples. Used for generating training data with specific seeds to alter during training",
|
||||
default=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--seed_lock",
|
||||
action="store_true",
|
||||
help="Locks the seed to the latent images so the same latent will always have the same noise",
|
||||
default=False
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = setup_parser()
|
||||
|
||||
args = parser.parse_args()
|
||||
args = train_util.read_config_from_file(args, parser)
|
||||
|
||||
train(args)
|
||||
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)
|
||||
130
testing/generate_lora_mapping.py
Normal file
130
testing/generate_lora_mapping.py
Normal file
@@ -0,0 +1,130 @@
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
import argparse
|
||||
import os
|
||||
import json
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
keymap_path = os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', 'stable_diffusion_sdxl.json')
|
||||
|
||||
# load keymap
|
||||
with open(keymap_path, 'r') as f:
|
||||
keymap = json.load(f)
|
||||
|
||||
lora_keymap = OrderedDict()
|
||||
|
||||
# convert keymap to lora key naming
|
||||
for ldm_key, diffusers_key in keymap['ldm_diffusers_keymap'].items():
|
||||
if ldm_key.endswith('.bias') or diffusers_key.endswith('.bias'):
|
||||
# skip it
|
||||
continue
|
||||
# sdxl has same te for locon with kohya and ours
|
||||
if ldm_key.startswith('conditioner'):
|
||||
#skip it
|
||||
continue
|
||||
# ignore vae
|
||||
if ldm_key.startswith('first_stage_model'):
|
||||
continue
|
||||
ldm_key = ldm_key.replace('model.diffusion_model.', 'lora_unet_')
|
||||
ldm_key = ldm_key.replace('.weight', '')
|
||||
ldm_key = ldm_key.replace('.', '_')
|
||||
|
||||
diffusers_key = diffusers_key.replace('unet_', 'lora_unet_')
|
||||
diffusers_key = diffusers_key.replace('.weight', '')
|
||||
diffusers_key = diffusers_key.replace('.', '_')
|
||||
|
||||
lora_keymap[f"{ldm_key}.alpha"] = f"{diffusers_key}.alpha"
|
||||
lora_keymap[f"{ldm_key}.lora_down.weight"] = f"{diffusers_key}.lora_down.weight"
|
||||
lora_keymap[f"{ldm_key}.lora_up.weight"] = f"{diffusers_key}.lora_up.weight"
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("input", help="input file")
|
||||
parser.add_argument("input2", help="input2 file")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# name = args.name
|
||||
# if args.sdxl:
|
||||
# name += '_sdxl'
|
||||
# elif args.sd2:
|
||||
# name += '_sd2'
|
||||
# else:
|
||||
# name += '_sd1'
|
||||
name = 'stable_diffusion_locon_sdxl'
|
||||
|
||||
locon_save = load_file(args.input)
|
||||
our_save = load_file(args.input2)
|
||||
|
||||
our_extra_keys = list(set(our_save.keys()) - set(locon_save.keys()))
|
||||
locon_extra_keys = list(set(locon_save.keys()) - set(our_save.keys()))
|
||||
|
||||
print(f"we have {len(our_extra_keys)} extra keys")
|
||||
print(f"locon has {len(locon_extra_keys)} extra keys")
|
||||
|
||||
save_dtype = torch.float16
|
||||
print(f"our extra keys: {our_extra_keys}")
|
||||
print(f"locon extra keys: {locon_extra_keys}")
|
||||
|
||||
|
||||
def export_state_dict(our_save):
|
||||
converted_state_dict = OrderedDict()
|
||||
for key, value in our_save.items():
|
||||
# test encoders share keys for some reason
|
||||
if key.startswith('lora_te'):
|
||||
converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
|
||||
else:
|
||||
converted_key = key
|
||||
for ldm_key, diffusers_key in lora_keymap.items():
|
||||
if converted_key == diffusers_key:
|
||||
converted_key = ldm_key
|
||||
|
||||
converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
|
||||
return converted_state_dict
|
||||
|
||||
def import_state_dict(loaded_state_dict):
|
||||
converted_state_dict = OrderedDict()
|
||||
for key, value in loaded_state_dict.items():
|
||||
if key.startswith('lora_te'):
|
||||
converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
|
||||
else:
|
||||
converted_key = key
|
||||
for ldm_key, diffusers_key in lora_keymap.items():
|
||||
if converted_key == ldm_key:
|
||||
converted_key = diffusers_key
|
||||
|
||||
converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
|
||||
return converted_state_dict
|
||||
|
||||
|
||||
# check it again
|
||||
converted_state_dict = export_state_dict(our_save)
|
||||
converted_extra_keys = list(set(converted_state_dict.keys()) - set(locon_save.keys()))
|
||||
locon_extra_keys = list(set(locon_save.keys()) - set(converted_state_dict.keys()))
|
||||
|
||||
|
||||
print(f"we have {len(converted_extra_keys)} extra keys")
|
||||
print(f"locon has {len(locon_extra_keys)} extra keys")
|
||||
|
||||
print(f"our extra keys: {converted_extra_keys}")
|
||||
|
||||
# convert back
|
||||
cycle_state_dict = import_state_dict(converted_state_dict)
|
||||
cycle_extra_keys = list(set(cycle_state_dict.keys()) - set(our_save.keys()))
|
||||
our_extra_keys = list(set(our_save.keys()) - set(cycle_state_dict.keys()))
|
||||
|
||||
print(f"we have {len(our_extra_keys)} extra keys")
|
||||
print(f"cycle has {len(cycle_extra_keys)} extra keys")
|
||||
|
||||
# save keymap
|
||||
to_save = OrderedDict()
|
||||
to_save['ldm_diffusers_keymap'] = lora_keymap
|
||||
|
||||
with open(os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', f'{name}.json'), 'w') as f:
|
||||
json.dump(to_save, f, indent=4)
|
||||
|
||||
|
||||
|
||||
463
testing/generate_weight_mappings.py
Normal file
463
testing/generate_weight_mappings.py
Normal file
@@ -0,0 +1,463 @@
|
||||
import argparse
|
||||
import gc
|
||||
import os
|
||||
import re
|
||||
import os
|
||||
# add project root to sys path
|
||||
import sys
|
||||
|
||||
from diffusers import DiffusionPipeline, StableDiffusionXLPipeline
|
||||
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import torch
|
||||
from diffusers.loaders import LoraLoaderMixin
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
KEYMAPS_FOLDER = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), 'toolkit', 'keymaps')
|
||||
|
||||
device = torch.device('cpu')
|
||||
dtype = torch.float32
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def get_reduced_shape(shape_tuple):
|
||||
# iterate though shape anr remove 1s
|
||||
new_shape = []
|
||||
for dim in shape_tuple:
|
||||
if dim != 1:
|
||||
new_shape.append(dim)
|
||||
return tuple(new_shape)
|
||||
|
||||
|
||||
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('--name', type=str, default='stable_diffusion', help='name for mapping to make')
|
||||
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
|
||||
parser.add_argument('--refiner', action='store_true', help='is refiner model')
|
||||
parser.add_argument('--ssd', action='store_true', help='is ssd model')
|
||||
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
file_path = args.file_1[0]
|
||||
|
||||
find_matches = False
|
||||
|
||||
print(f'Loading diffusers model')
|
||||
|
||||
ignore_ldm_begins_with = []
|
||||
|
||||
diffusers_file_path = file_path
|
||||
if args.ssd:
|
||||
diffusers_file_path = "segmind/SSD-1B"
|
||||
|
||||
# if args.refiner:
|
||||
# diffusers_file_path = "stabilityai/stable-diffusion-xl-refiner-1.0"
|
||||
|
||||
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
|
||||
|
||||
if not args.refiner:
|
||||
|
||||
diffusers_model_config = ModelConfig(
|
||||
name_or_path=diffusers_file_path,
|
||||
is_xl=args.sdxl,
|
||||
is_v2=args.sd2,
|
||||
is_ssd=args.ssd,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd = StableDiffusion(
|
||||
model_config=diffusers_model_config,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd.load_model()
|
||||
# delete things we dont need
|
||||
del diffusers_sd.tokenizer
|
||||
flush()
|
||||
|
||||
print(f'Loading ldm model')
|
||||
diffusers_state_dict = diffusers_sd.state_dict()
|
||||
else:
|
||||
# refiner wont work directly with stable diffusion
|
||||
# so we need to load the model and then load the state dict
|
||||
diffusers_pipeline = StableDiffusionXLPipeline.from_single_file(
|
||||
diffusers_file_path,
|
||||
torch_dtype=torch.float16,
|
||||
use_safetensors=True,
|
||||
variant="fp16",
|
||||
).to(device)
|
||||
# diffusers_pipeline = StableDiffusionXLPipeline.from_single_file(
|
||||
# file_path,
|
||||
# torch_dtype=torch.float16,
|
||||
# use_safetensors=True,
|
||||
# variant="fp16",
|
||||
# ).to(device)
|
||||
|
||||
SD_PREFIX_VAE = "vae"
|
||||
SD_PREFIX_UNET = "unet"
|
||||
SD_PREFIX_REFINER_UNET = "refiner_unet"
|
||||
SD_PREFIX_TEXT_ENCODER = "te"
|
||||
|
||||
SD_PREFIX_TEXT_ENCODER1 = "te0"
|
||||
SD_PREFIX_TEXT_ENCODER2 = "te1"
|
||||
|
||||
diffusers_state_dict = OrderedDict()
|
||||
for k, v in diffusers_pipeline.vae.state_dict().items():
|
||||
new_key = k if k.startswith(f"{SD_PREFIX_VAE}") else f"{SD_PREFIX_VAE}_{k}"
|
||||
diffusers_state_dict[new_key] = v
|
||||
for k, v in diffusers_pipeline.text_encoder_2.state_dict().items():
|
||||
new_key = k if k.startswith(f"{SD_PREFIX_TEXT_ENCODER2}_") else f"{SD_PREFIX_TEXT_ENCODER2}_{k}"
|
||||
diffusers_state_dict[new_key] = v
|
||||
for k, v in diffusers_pipeline.unet.state_dict().items():
|
||||
new_key = k if k.startswith(f"{SD_PREFIX_UNET}_") else f"{SD_PREFIX_UNET}_{k}"
|
||||
diffusers_state_dict[new_key] = v
|
||||
|
||||
# add ignore ones as we are only going to focus on unet and copy the rest
|
||||
# ignore_ldm_begins_with = ["conditioner.", "first_stage_model."]
|
||||
|
||||
diffusers_dict_keys = list(diffusers_state_dict.keys())
|
||||
|
||||
ldm_state_dict = load_file(file_path)
|
||||
ldm_dict_keys = list(ldm_state_dict.keys())
|
||||
|
||||
ldm_diffusers_keymap = OrderedDict()
|
||||
ldm_diffusers_shape_map = OrderedDict()
|
||||
ldm_operator_map = OrderedDict()
|
||||
diffusers_operator_map = OrderedDict()
|
||||
|
||||
total_keys = len(ldm_dict_keys)
|
||||
|
||||
matched_ldm_keys = []
|
||||
matched_diffusers_keys = []
|
||||
|
||||
error_margin = 1e-8
|
||||
|
||||
tmp_merge_key = "TMP___MERGE"
|
||||
|
||||
te_suffix = ''
|
||||
proj_pattern_weight = None
|
||||
proj_pattern_bias = None
|
||||
text_proj_layer = None
|
||||
if args.sdxl or args.ssd:
|
||||
te_suffix = '1'
|
||||
ldm_res_block_prefix = "conditioner.embedders.1.model.transformer.resblocks"
|
||||
proj_pattern_weight = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
|
||||
proj_pattern_bias = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
|
||||
text_proj_layer = "conditioner.embedders.1.model.text_projection"
|
||||
if args.refiner:
|
||||
te_suffix = '1'
|
||||
ldm_res_block_prefix = "conditioner.embedders.0.model.transformer.resblocks"
|
||||
proj_pattern_weight = r"conditioner\.embedders\.0\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
|
||||
proj_pattern_bias = r"conditioner\.embedders\.0\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
|
||||
text_proj_layer = "conditioner.embedders.0.model.text_projection"
|
||||
if args.sd2:
|
||||
te_suffix = ''
|
||||
ldm_res_block_prefix = "cond_stage_model.model.transformer.resblocks"
|
||||
proj_pattern_weight = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
|
||||
proj_pattern_bias = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
|
||||
text_proj_layer = "cond_stage_model.model.text_projection"
|
||||
|
||||
if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
if "conditioner.embedders.1.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection"].shape[0])
|
||||
elif "conditioner.embedders.0.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.0.model.text_projection"].shape[0])
|
||||
else:
|
||||
d_model = 1024
|
||||
|
||||
# do pre known merging
|
||||
for ldm_key in ldm_dict_keys:
|
||||
try:
|
||||
match = re.match(proj_pattern_weight, ldm_key)
|
||||
if match:
|
||||
number = int(match.group(1))
|
||||
new_val = torch.cat([
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"],
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight"],
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight"],
|
||||
], dim=0)
|
||||
# add to matched so we dont check them
|
||||
matched_diffusers_keys.append(
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight")
|
||||
matched_diffusers_keys.append(
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight")
|
||||
matched_diffusers_keys.append(
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight")
|
||||
# make diffusers convertable_dict
|
||||
diffusers_state_dict[
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.{tmp_merge_key}.weight"] = new_val
|
||||
|
||||
# add operator
|
||||
ldm_operator_map[ldm_key] = {
|
||||
"cat": [
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight",
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight",
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight",
|
||||
],
|
||||
}
|
||||
|
||||
# text_model_dict[new_key + ".q_proj.weight"] = checkpoint[key][:d_model, :]
|
||||
# text_model_dict[new_key + ".k_proj.weight"] = checkpoint[key][d_model: d_model * 2, :]
|
||||
# text_model_dict[new_key + ".v_proj.weight"] = checkpoint[key][d_model * 2:, :]
|
||||
|
||||
# add diffusers operators
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
|
||||
f"0:{d_model}, :"
|
||||
]
|
||||
}
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
|
||||
f"{d_model}:{d_model * 2}, :"
|
||||
]
|
||||
}
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
|
||||
f"{d_model * 2}:, :"
|
||||
]
|
||||
}
|
||||
|
||||
match = re.match(proj_pattern_bias, ldm_key)
|
||||
if match:
|
||||
number = int(match.group(1))
|
||||
new_val = torch.cat([
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias"],
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias"],
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias"],
|
||||
], dim=0)
|
||||
# add to matched so we dont check them
|
||||
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias")
|
||||
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias")
|
||||
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias")
|
||||
# make diffusers convertable_dict
|
||||
diffusers_state_dict[
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.{tmp_merge_key}.bias"] = new_val
|
||||
|
||||
# add operator
|
||||
ldm_operator_map[ldm_key] = {
|
||||
"cat": [
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias",
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias",
|
||||
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias",
|
||||
],
|
||||
}
|
||||
|
||||
# add diffusers operators
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
|
||||
f"0:{d_model}, :"
|
||||
]
|
||||
}
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
|
||||
f"{d_model}:{d_model * 2}, :"
|
||||
]
|
||||
}
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias"] = {
|
||||
"slice": [
|
||||
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
|
||||
f"{d_model * 2}:, :"
|
||||
]
|
||||
}
|
||||
except Exception as e:
|
||||
print(f"Error on key {ldm_key}")
|
||||
print(e)
|
||||
|
||||
# update keys
|
||||
diffusers_dict_keys = list(diffusers_state_dict.keys())
|
||||
|
||||
pbar = tqdm(ldm_dict_keys, desc='Matching ldm-diffusers keys', total=total_keys)
|
||||
# run through all weights and check mse between them to find matches
|
||||
for ldm_key in ldm_dict_keys:
|
||||
ldm_shape_tuple = ldm_state_dict[ldm_key].shape
|
||||
ldm_reduced_shape_tuple = get_reduced_shape(ldm_shape_tuple)
|
||||
for diffusers_key in diffusers_dict_keys:
|
||||
diffusers_shape_tuple = diffusers_state_dict[diffusers_key].shape
|
||||
diffusers_reduced_shape_tuple = get_reduced_shape(diffusers_shape_tuple)
|
||||
|
||||
# That was easy. Same key
|
||||
# if ldm_key == diffusers_key:
|
||||
# ldm_diffusers_keymap[ldm_key] = diffusers_key
|
||||
# matched_ldm_keys.append(ldm_key)
|
||||
# matched_diffusers_keys.append(diffusers_key)
|
||||
# break
|
||||
|
||||
# if we already have this key mapped, skip it
|
||||
if diffusers_key in matched_diffusers_keys:
|
||||
continue
|
||||
|
||||
# if reduced shapes do not match skip it
|
||||
if ldm_reduced_shape_tuple != diffusers_reduced_shape_tuple:
|
||||
continue
|
||||
|
||||
ldm_weight = ldm_state_dict[ldm_key]
|
||||
did_reduce_ldm = False
|
||||
diffusers_weight = diffusers_state_dict[diffusers_key]
|
||||
did_reduce_diffusers = False
|
||||
|
||||
# reduce the shapes to match if they are not the same
|
||||
if ldm_shape_tuple != ldm_reduced_shape_tuple:
|
||||
ldm_weight = ldm_weight.view(ldm_reduced_shape_tuple)
|
||||
did_reduce_ldm = True
|
||||
|
||||
if diffusers_shape_tuple != diffusers_reduced_shape_tuple:
|
||||
diffusers_weight = diffusers_weight.view(diffusers_reduced_shape_tuple)
|
||||
did_reduce_diffusers = True
|
||||
|
||||
# check to see if they match within a margin of error
|
||||
mse = torch.nn.functional.mse_loss(ldm_weight.float(), diffusers_weight.float())
|
||||
if mse < error_margin:
|
||||
ldm_diffusers_keymap[ldm_key] = diffusers_key
|
||||
matched_ldm_keys.append(ldm_key)
|
||||
matched_diffusers_keys.append(diffusers_key)
|
||||
|
||||
if did_reduce_ldm or did_reduce_diffusers:
|
||||
ldm_diffusers_shape_map[ldm_key] = (ldm_shape_tuple, diffusers_shape_tuple)
|
||||
if did_reduce_ldm:
|
||||
del ldm_weight
|
||||
if did_reduce_diffusers:
|
||||
del diffusers_weight
|
||||
flush()
|
||||
|
||||
break
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
pbar.close()
|
||||
|
||||
name = args.name
|
||||
if args.sdxl:
|
||||
name += '_sdxl'
|
||||
elif args.ssd:
|
||||
name += '_ssd'
|
||||
elif args.refiner:
|
||||
name += '_refiner'
|
||||
elif args.sd2:
|
||||
name += '_sd2'
|
||||
else:
|
||||
name += '_sd1'
|
||||
|
||||
# if len(matched_ldm_keys) != len(matched_diffusers_keys):
|
||||
unmatched_ldm_keys = [x for x in ldm_dict_keys if x not in matched_ldm_keys]
|
||||
unmatched_diffusers_keys = [x for x in diffusers_dict_keys if x not in matched_diffusers_keys]
|
||||
# has unmatched keys
|
||||
|
||||
has_unmatched_keys = len(unmatched_ldm_keys) > 0 or len(unmatched_diffusers_keys) > 0
|
||||
|
||||
|
||||
def get_slices_from_string(s: str) -> tuple:
|
||||
slice_strings = s.split(',')
|
||||
slices = [eval(f"slice({component.strip()})") for component in slice_strings]
|
||||
return tuple(slices)
|
||||
|
||||
|
||||
if has_unmatched_keys:
|
||||
|
||||
print(
|
||||
f"Found {len(unmatched_ldm_keys)} unmatched ldm keys and {len(unmatched_diffusers_keys)} unmatched diffusers keys")
|
||||
|
||||
unmatched_obj = OrderedDict()
|
||||
unmatched_obj['ldm'] = OrderedDict()
|
||||
unmatched_obj['diffusers'] = OrderedDict()
|
||||
|
||||
print(f"Gathering info on unmatched keys")
|
||||
|
||||
for key in tqdm(unmatched_ldm_keys, desc='Unmatched LDM keys'):
|
||||
# get min, max, mean, std
|
||||
weight = ldm_state_dict[key]
|
||||
weight_min = weight.min().item()
|
||||
weight_max = weight.max().item()
|
||||
unmatched_obj['ldm'][key] = {
|
||||
'shape': weight.shape,
|
||||
"min": weight_min,
|
||||
"max": weight_max,
|
||||
}
|
||||
del weight
|
||||
flush()
|
||||
|
||||
for key in tqdm(unmatched_diffusers_keys, desc='Unmatched Diffusers keys'):
|
||||
# get min, max, mean, std
|
||||
weight = diffusers_state_dict[key]
|
||||
weight_min = weight.min().item()
|
||||
weight_max = weight.max().item()
|
||||
unmatched_obj['diffusers'][key] = {
|
||||
"shape": weight.shape,
|
||||
"min": weight_min,
|
||||
"max": weight_max,
|
||||
}
|
||||
del weight
|
||||
flush()
|
||||
|
||||
unmatched_path = os.path.join(KEYMAPS_FOLDER, f'{name}_unmatched.json')
|
||||
with open(unmatched_path, 'w') as f:
|
||||
f.write(json.dumps(unmatched_obj, indent=4))
|
||||
|
||||
print(f'Saved unmatched keys to {unmatched_path}')
|
||||
|
||||
# save ldm remainders
|
||||
remaining_ldm_values = OrderedDict()
|
||||
for key in unmatched_ldm_keys:
|
||||
remaining_ldm_values[key] = ldm_state_dict[key].detach().to('cpu', torch.float16)
|
||||
|
||||
save_file(remaining_ldm_values, os.path.join(KEYMAPS_FOLDER, f'{name}_ldm_base.safetensors'))
|
||||
print(f'Saved remaining ldm values to {os.path.join(KEYMAPS_FOLDER, f"{name}_ldm_base.safetensors")}')
|
||||
|
||||
# do cleanup of some left overs and bugs
|
||||
to_remove = []
|
||||
for ldm_key, diffusers_key in ldm_diffusers_keymap.items():
|
||||
# get rid of tmp merge keys used to slicing
|
||||
if tmp_merge_key in diffusers_key or tmp_merge_key in ldm_key:
|
||||
to_remove.append(ldm_key)
|
||||
|
||||
for key in to_remove:
|
||||
del ldm_diffusers_keymap[key]
|
||||
|
||||
to_remove = []
|
||||
# remove identical shape mappings. Not sure why they exist but they do
|
||||
for ldm_key, shape_list in ldm_diffusers_shape_map.items():
|
||||
# remove identical shape mappings. Not sure why they exist but they do
|
||||
# convert to json string to make it easier to compare
|
||||
ldm_shape = json.dumps(shape_list[0])
|
||||
diffusers_shape = json.dumps(shape_list[1])
|
||||
if ldm_shape == diffusers_shape:
|
||||
to_remove.append(ldm_key)
|
||||
|
||||
for key in to_remove:
|
||||
del ldm_diffusers_shape_map[key]
|
||||
|
||||
dest_path = os.path.join(KEYMAPS_FOLDER, f'{name}.json')
|
||||
save_obj = OrderedDict()
|
||||
save_obj["ldm_diffusers_keymap"] = ldm_diffusers_keymap
|
||||
save_obj["ldm_diffusers_shape_map"] = ldm_diffusers_shape_map
|
||||
save_obj["ldm_diffusers_operator_map"] = ldm_operator_map
|
||||
save_obj["diffusers_ldm_operator_map"] = diffusers_operator_map
|
||||
with open(dest_path, 'w') as f:
|
||||
f.write(json.dumps(save_obj, indent=4))
|
||||
|
||||
print(f'Saved keymap to {dest_path}')
|
||||
107
testing/test_bucket_dataloader.py
Normal file
107
testing/test_bucket_dataloader.py
Normal file
@@ -0,0 +1,107 @@
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchvision import transforms
|
||||
import sys
|
||||
import os
|
||||
import cv2
|
||||
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
from toolkit.image_utils import show_img
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from library.model_util import load_vae
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.data_loader import AiToolkitDataset, get_dataloader_from_datasets, \
|
||||
trigger_dataloader_setup_epoch
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('dataset_folder', type=str, default='input')
|
||||
parser.add_argument('--epochs', type=int, default=1)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
dataset_folder = args.dataset_folder
|
||||
resolution = 1024
|
||||
bucket_tolerance = 64
|
||||
batch_size = 1
|
||||
|
||||
|
||||
##
|
||||
|
||||
dataset_config = DatasetConfig(
|
||||
dataset_path=dataset_folder,
|
||||
resolution=resolution,
|
||||
caption_ext='json',
|
||||
default_caption='default',
|
||||
buckets=True,
|
||||
bucket_tolerance=bucket_tolerance,
|
||||
poi='person',
|
||||
augmentations=[
|
||||
{
|
||||
'method': 'RandomBrightnessContrast',
|
||||
'brightness_limit': (-0.3, 0.3),
|
||||
'contrast_limit': (-0.3, 0.3),
|
||||
'brightness_by_max': False,
|
||||
'p': 1.0
|
||||
},
|
||||
{
|
||||
'method': 'HueSaturationValue',
|
||||
'hue_shift_limit': (-0, 0),
|
||||
'sat_shift_limit': (-40, 40),
|
||||
'val_shift_limit': (-40, 40),
|
||||
'p': 1.0
|
||||
},
|
||||
# {
|
||||
# 'method': 'RGBShift',
|
||||
# 'r_shift_limit': (-20, 20),
|
||||
# 'g_shift_limit': (-20, 20),
|
||||
# 'b_shift_limit': (-20, 20),
|
||||
# 'p': 1.0
|
||||
# },
|
||||
]
|
||||
|
||||
|
||||
)
|
||||
|
||||
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size)
|
||||
|
||||
|
||||
# run through an epoch ang check sizes
|
||||
dataloader_iterator = iter(dataloader)
|
||||
for epoch in range(args.epochs):
|
||||
for batch in dataloader:
|
||||
batch: 'DataLoaderBatchDTO'
|
||||
img_batch = batch.tensor
|
||||
|
||||
chunks = torch.chunk(img_batch, batch_size, dim=0)
|
||||
# put them so they are size by side
|
||||
big_img = torch.cat(chunks, dim=3)
|
||||
big_img = big_img.squeeze(0)
|
||||
|
||||
min_val = big_img.min()
|
||||
max_val = big_img.max()
|
||||
|
||||
big_img = (big_img / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
# convert to image
|
||||
img = transforms.ToPILImage()(big_img)
|
||||
|
||||
show_img(img)
|
||||
|
||||
time.sleep(1.0)
|
||||
# if not last epoch
|
||||
if epoch < args.epochs - 1:
|
||||
trigger_dataloader_setup_epoch(dataloader)
|
||||
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
print('done')
|
||||
172
testing/test_model_load_save.py
Normal file
172
testing/test_model_load_save.py
Normal file
@@ -0,0 +1,172 @@
|
||||
import argparse
|
||||
import os
|
||||
# add project root to sys path
|
||||
import sys
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import torch
|
||||
from diffusers.loaders import LoraLoaderMixin
|
||||
from safetensors.torch import load_file
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
|
||||
from toolkit.config_modules import ModelConfig
|
||||
from toolkit.paths import KEYMAPS_ROOT
|
||||
from toolkit.saving import convert_state_dict_to_ldm_with_mapping, get_ldm_state_dict_from_diffusers
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
# 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
|
||||
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
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 an LDM model'
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'--is_xl',
|
||||
action='store_true',
|
||||
help='Is the model an XL model'
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'--is_v2',
|
||||
action='store_true',
|
||||
help='Is the model a v2 model'
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
find_matches = False
|
||||
|
||||
print("Loading model")
|
||||
state_dict_file_1 = load_file(args.file_1[0])
|
||||
state_dict_1_keys = list(state_dict_file_1.keys())
|
||||
|
||||
print("Loading model into diffusers format")
|
||||
model_config = ModelConfig(
|
||||
name_or_path=args.file_1[0],
|
||||
is_xl=args.is_xl
|
||||
)
|
||||
sd = StableDiffusion(
|
||||
model_config=model_config,
|
||||
device=device,
|
||||
)
|
||||
sd.load_model()
|
||||
|
||||
# load our base
|
||||
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl_ldm_base.safetensors')
|
||||
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl.json')
|
||||
|
||||
print("Converting model back to LDM")
|
||||
version_string = '1'
|
||||
if args.is_v2:
|
||||
version_string = '2'
|
||||
if args.is_xl:
|
||||
version_string = 'sdxl'
|
||||
# convert the state dict
|
||||
state_dict_file_2 = get_ldm_state_dict_from_diffusers(
|
||||
sd.state_dict(),
|
||||
version_string,
|
||||
device='cpu',
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
# 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()
|
||||
|
||||
if len(keys_not_in_state_dict_2) == 0 and len(keys_not_in_state_dict_1) == 0:
|
||||
print("All keys match!")
|
||||
print("Checking values...")
|
||||
mismatch_keys = []
|
||||
loss = torch.nn.MSELoss()
|
||||
tolerance = 1e-6
|
||||
for key in tqdm(keys_in_both):
|
||||
if loss(state_dict_file_1[key], state_dict_file_2[key]) > tolerance:
|
||||
print(f"Values for key {key} don't match!")
|
||||
print(f"Loss: {loss(state_dict_file_1[key], state_dict_file_2[key])}")
|
||||
mismatch_keys.append(key)
|
||||
|
||||
if len(mismatch_keys) == 0:
|
||||
print("All values match!")
|
||||
else:
|
||||
print("Some valued font match!")
|
||||
print(mismatch_keys)
|
||||
mismatched_path = os.path.join(project_root, 'config', 'mismatch.json')
|
||||
with open(mismatched_path, 'w') as f:
|
||||
f.write(json.dumps(mismatch_keys, indent=4))
|
||||
exit(0)
|
||||
|
||||
else:
|
||||
print("Keys don't match!, generating info...")
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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_1_filename}_loop.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)
|
||||
50
toolkit/basic.py
Normal file
50
toolkit/basic.py
Normal file
@@ -0,0 +1,50 @@
|
||||
import gc
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def flush(garbage_collect=True):
|
||||
torch.cuda.empty_cache()
|
||||
if garbage_collect:
|
||||
gc.collect()
|
||||
|
||||
|
||||
def get_mean_std(tensor):
|
||||
if len(tensor.shape) == 3:
|
||||
tensor = tensor.unsqueeze(0)
|
||||
elif len(tensor.shape) != 4:
|
||||
raise Exception("Expected tensor of shape (batch_size, channels, width, height)")
|
||||
mean, variance = torch.mean(
|
||||
tensor, dim=[2, 3], keepdim=True
|
||||
), torch.var(
|
||||
tensor, dim=[2, 3],
|
||||
keepdim=True
|
||||
)
|
||||
std = torch.sqrt(variance + 1e-5)
|
||||
return mean, std
|
||||
|
||||
|
||||
def adain(content_features, style_features):
|
||||
# Assumes that the content and style features are of shape (batch_size, channels, width, height)
|
||||
|
||||
# Step 1: Calculate mean and variance of content features
|
||||
content_mean, content_var = torch.mean(content_features, dim=[2, 3], keepdim=True), torch.var(content_features,
|
||||
dim=[2, 3],
|
||||
keepdim=True)
|
||||
# Step 2: Calculate mean and variance of style features
|
||||
style_mean, style_var = torch.mean(style_features, dim=[2, 3], keepdim=True), torch.var(style_features, dim=[2, 3],
|
||||
keepdim=True)
|
||||
|
||||
# Step 3: Normalize content features
|
||||
content_std = torch.sqrt(content_var + 1e-5)
|
||||
normalized_content = (content_features - content_mean) / content_std
|
||||
|
||||
# Step 4: Scale and shift normalized content with style's statistics
|
||||
style_std = torch.sqrt(style_var + 1e-5)
|
||||
stylized_content = normalized_content * style_std + style_mean
|
||||
|
||||
return stylized_content
|
||||
127
toolkit/buckets.py
Normal file
127
toolkit/buckets.py
Normal file
@@ -0,0 +1,127 @@
|
||||
from typing import Type, List, Union, TypedDict
|
||||
|
||||
|
||||
class BucketResolution(TypedDict):
|
||||
width: int
|
||||
height: int
|
||||
|
||||
|
||||
# resolutions SDXL was trained on with a 1024x1024 base resolution
|
||||
resolutions_1024: List[BucketResolution] = [
|
||||
# SDXL Base resolution
|
||||
{"width": 1024, "height": 1024},
|
||||
# SDXL Resolutions, widescreen
|
||||
{"width": 2048, "height": 512},
|
||||
{"width": 1984, "height": 512},
|
||||
{"width": 1920, "height": 512},
|
||||
{"width": 1856, "height": 512},
|
||||
{"width": 1792, "height": 576},
|
||||
{"width": 1728, "height": 576},
|
||||
{"width": 1664, "height": 576},
|
||||
{"width": 1600, "height": 640},
|
||||
{"width": 1536, "height": 640},
|
||||
{"width": 1472, "height": 704},
|
||||
{"width": 1408, "height": 704},
|
||||
{"width": 1344, "height": 704},
|
||||
{"width": 1344, "height": 768},
|
||||
{"width": 1280, "height": 768},
|
||||
{"width": 1216, "height": 832},
|
||||
{"width": 1152, "height": 832},
|
||||
{"width": 1152, "height": 896},
|
||||
{"width": 1088, "height": 896},
|
||||
{"width": 1088, "height": 960},
|
||||
{"width": 1024, "height": 960},
|
||||
# SDXL Resolutions, portrait
|
||||
{"width": 960, "height": 1024},
|
||||
{"width": 960, "height": 1088},
|
||||
{"width": 896, "height": 1088},
|
||||
{"width": 896, "height": 1152}, # 2:3
|
||||
{"width": 832, "height": 1152},
|
||||
{"width": 832, "height": 1216},
|
||||
{"width": 768, "height": 1280},
|
||||
{"width": 768, "height": 1344},
|
||||
{"width": 704, "height": 1408},
|
||||
{"width": 704, "height": 1472},
|
||||
{"width": 640, "height": 1536},
|
||||
{"width": 640, "height": 1600},
|
||||
{"width": 576, "height": 1664},
|
||||
{"width": 576, "height": 1728},
|
||||
{"width": 576, "height": 1792},
|
||||
{"width": 512, "height": 1856},
|
||||
{"width": 512, "height": 1920},
|
||||
{"width": 512, "height": 1984},
|
||||
{"width": 512, "height": 2048},
|
||||
]
|
||||
|
||||
|
||||
def get_bucket_sizes(resolution: int = 512, divisibility: int = 8) -> List[BucketResolution]:
|
||||
# determine scaler form 1024 to resolution
|
||||
scaler = resolution / 1024
|
||||
|
||||
bucket_size_list = []
|
||||
for bucket in resolutions_1024:
|
||||
# must be divisible by 8
|
||||
width = int(bucket["width"] * scaler)
|
||||
height = int(bucket["height"] * scaler)
|
||||
if width % divisibility != 0:
|
||||
width = width - (width % divisibility)
|
||||
if height % divisibility != 0:
|
||||
height = height - (height % divisibility)
|
||||
bucket_size_list.append({"width": width, "height": height})
|
||||
|
||||
return bucket_size_list
|
||||
|
||||
|
||||
def get_resolution(width, height):
|
||||
num_pixels = width * height
|
||||
# determine same number of pixels for square image
|
||||
square_resolution = int(num_pixels ** 0.5)
|
||||
return square_resolution
|
||||
|
||||
|
||||
def get_bucket_for_image_size(
|
||||
width: int,
|
||||
height: int,
|
||||
bucket_size_list: List[BucketResolution] = None,
|
||||
resolution: Union[int, None] = None,
|
||||
divisibility: int = 8
|
||||
) -> BucketResolution:
|
||||
|
||||
if bucket_size_list is None and resolution is None:
|
||||
# get resolution from width and height
|
||||
resolution = get_resolution(width, height)
|
||||
if bucket_size_list is None:
|
||||
# if real resolution is smaller, use that instead
|
||||
real_resolution = get_resolution(width, height)
|
||||
resolution = min(resolution, real_resolution)
|
||||
bucket_size_list = get_bucket_sizes(resolution=resolution, divisibility=divisibility)
|
||||
|
||||
# Check for exact match first
|
||||
for bucket in bucket_size_list:
|
||||
if bucket["width"] == width and bucket["height"] == height:
|
||||
return bucket
|
||||
|
||||
# If exact match not found, find the closest bucket
|
||||
closest_bucket = None
|
||||
min_removed_pixels = float("inf")
|
||||
|
||||
for bucket in bucket_size_list:
|
||||
scale_w = bucket["width"] / width
|
||||
scale_h = bucket["height"] / height
|
||||
|
||||
# To minimize pixels, we use the larger scale factor to minimize the amount that has to be cropped.
|
||||
scale = max(scale_w, scale_h)
|
||||
|
||||
new_width = int(width * scale)
|
||||
new_height = int(height * scale)
|
||||
|
||||
removed_pixels = (new_width - bucket["width"]) * new_height + (new_height - bucket["height"]) * new_width
|
||||
|
||||
if removed_pixels < min_removed_pixels:
|
||||
min_removed_pixels = removed_pixels
|
||||
closest_bucket = bucket
|
||||
|
||||
if closest_bucket is None:
|
||||
raise ValueError("No suitable bucket found")
|
||||
|
||||
return closest_bucket
|
||||
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,42 @@ def get_cwd_abs_path(path):
|
||||
return path
|
||||
|
||||
|
||||
def preprocess_config(config: OrderedDict):
|
||||
def replace_env_vars_in_string(s: str) -> str:
|
||||
"""
|
||||
Replace placeholders like ${VAR_NAME} with the value of the corresponding environment variable.
|
||||
If the environment variable is not set, raise an error.
|
||||
"""
|
||||
|
||||
def replacer(match):
|
||||
var_name = match.group(1)
|
||||
value = os.environ.get(var_name)
|
||||
|
||||
if value is None:
|
||||
raise ValueError(f"Environment variable {var_name} not set. Please ensure it's defined before proceeding.")
|
||||
|
||||
return value
|
||||
|
||||
return re.sub(r'\$\{([^}]+)\}', replacer, s)
|
||||
|
||||
|
||||
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 +66,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
|
||||
@@ -66,13 +99,14 @@ def get_config(config_file_path):
|
||||
raise ValueError(f"Could not find config file {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:
|
||||
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:
|
||||
config = yaml.load(f, Loader=fixed_loader)
|
||||
else:
|
||||
raise ValueError(f"Config file {config_file_path} must be a json or yaml file")
|
||||
with open(real_config_path, 'r', encoding='utf-8') as f:
|
||||
content = f.read()
|
||||
content_with_env_replaced = replace_env_vars_in_string(content)
|
||||
if real_config_path.endswith('.json') or real_config_path.endswith('.jsonc'):
|
||||
config = json.loads(content_with_env_replaced, object_pairs_hook=OrderedDict)
|
||||
elif real_config_path.endswith('.yaml') or real_config_path.endswith('.yml'):
|
||||
config = yaml.load(content_with_env_replaced, 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,15 @@
|
||||
from typing import List
|
||||
import os
|
||||
import time
|
||||
from typing import List, Optional, Literal, Union
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
|
||||
ImgExt = Literal['jpg', 'png', 'webp']
|
||||
|
||||
SaveFormat = Literal['safetensors', 'diffusers']
|
||||
|
||||
|
||||
class SaveConfig:
|
||||
@@ -6,6 +17,9 @@ class SaveConfig:
|
||||
self.save_every: int = kwargs.get('save_every', 1000)
|
||||
self.dtype: str = kwargs.get('save_dtype', 'float16')
|
||||
self.max_step_saves_to_keep: int = kwargs.get('max_step_saves_to_keep', 5)
|
||||
self.save_format: SaveFormat = kwargs.get('save_format', 'safetensors')
|
||||
if self.save_format not in ['safetensors', 'diffusers']:
|
||||
raise ValueError(f"save_format must be safetensors or diffusers, got {self.save_format}")
|
||||
|
||||
|
||||
class LogingConfig:
|
||||
@@ -17,6 +31,7 @@ class LogingConfig:
|
||||
|
||||
class SampleConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.sampler: str = kwargs.get('sampler', 'ddpm')
|
||||
self.sample_every: int = kwargs.get('sample_every', 100)
|
||||
self.width: int = kwargs.get('width', 512)
|
||||
self.height: int = kwargs.get('height', 512)
|
||||
@@ -27,40 +42,201 @@ 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)
|
||||
self.ext: ImgExt = kwargs.get('format', 'jpg')
|
||||
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5) # step to start using refiner on sample if it exists
|
||||
|
||||
|
||||
class LormModuleSettingsConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.contains: str = kwargs.get('contains', '4nt$3')
|
||||
self.extract_mode: str = kwargs.get('extract_mode', 'ratio')
|
||||
# min num parameters to attach to
|
||||
self.parameter_threshold: int = kwargs.get('parameter_threshold', 0)
|
||||
self.extract_mode_param: dict = kwargs.get('extract_mode_param', 0.25)
|
||||
|
||||
|
||||
class LoRMConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.extract_mode: str = kwargs.get('extract_mode', 'ratio')
|
||||
self.do_conv: bool = kwargs.get('do_conv', False)
|
||||
self.extract_mode_param: dict = kwargs.get('extract_mode_param', 0.25)
|
||||
self.parameter_threshold: int = kwargs.get('parameter_threshold', 0)
|
||||
module_settings = kwargs.get('module_settings', [])
|
||||
default_module_settings = {
|
||||
'extract_mode': self.extract_mode,
|
||||
'extract_mode_param': self.extract_mode_param,
|
||||
'parameter_threshold': self.parameter_threshold,
|
||||
}
|
||||
module_settings = [{**default_module_settings, **module_setting, } for module_setting in module_settings]
|
||||
self.module_settings: List[LormModuleSettingsConfig] = [LormModuleSettingsConfig(**module_setting) for
|
||||
module_setting in module_settings]
|
||||
|
||||
def get_config_for_module(self, block_name):
|
||||
for setting in self.module_settings:
|
||||
contain_pieces = setting.contains.split('|')
|
||||
if all(contain_piece in block_name for contain_piece in contain_pieces):
|
||||
return setting
|
||||
# try replacing the . with _
|
||||
contain_pieces = setting.contains.replace('.', '_').split('|')
|
||||
if all(contain_piece in block_name for contain_piece in contain_pieces):
|
||||
return setting
|
||||
# do default
|
||||
return LormModuleSettingsConfig(**{
|
||||
'extract_mode': self.extract_mode,
|
||||
'extract_mode_param': self.extract_mode_param,
|
||||
'parameter_threshold': self.parameter_threshold,
|
||||
})
|
||||
|
||||
|
||||
NetworkType = Literal['lora', 'locon', 'lorm']
|
||||
|
||||
|
||||
class NetworkConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.type: str = kwargs.get('type', 'lora')
|
||||
self.type: NetworkType = kwargs.get('type', 'lora')
|
||||
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)
|
||||
self.dropout: Union[float, None] = kwargs.get('dropout', None)
|
||||
|
||||
self.lorm_config: Union[LoRMConfig, None] = None
|
||||
lorm = kwargs.get('lorm', None)
|
||||
if lorm is not None:
|
||||
self.lorm_config: LoRMConfig = LoRMConfig(**lorm)
|
||||
|
||||
if self.type == 'lorm':
|
||||
# set linear to arbitrary values so it makes them
|
||||
self.linear = 4
|
||||
self.rank = 4
|
||||
if self.lorm_config.do_conv:
|
||||
self.conv = 4
|
||||
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+']
|
||||
|
||||
|
||||
class AdapterConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip
|
||||
self.in_channels: int = kwargs.get('in_channels', 3)
|
||||
self.channels: List[int] = kwargs.get('channels', [320, 640, 1280, 1280])
|
||||
self.num_res_blocks: int = kwargs.get('num_res_blocks', 2)
|
||||
self.downscale_factor: int = kwargs.get('downscale_factor', 8)
|
||||
self.adapter_type: str = kwargs.get('adapter_type', 'full_adapter')
|
||||
self.image_dir: str = kwargs.get('image_dir', None)
|
||||
self.test_img_path: str = kwargs.get('test_img_path', None)
|
||||
self.train: str = kwargs.get('train', False)
|
||||
self.image_encoder_path: str = kwargs.get('image_encoder_path', None)
|
||||
self.name_or_path = kwargs.get('name_or_path', None)
|
||||
|
||||
|
||||
class EmbeddingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.trigger = kwargs.get('trigger', 'custom_embedding')
|
||||
self.tokens = kwargs.get('tokens', 4)
|
||||
self.init_words = kwargs.get('init_words', '*')
|
||||
self.save_format = kwargs.get('save_format', 'safetensors')
|
||||
|
||||
|
||||
ContentOrStyleType = Literal['balanced', 'style', 'content']
|
||||
LossTarget = Literal['noise', 'source', 'unaugmented', 'differential_noise']
|
||||
|
||||
|
||||
class TrainConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.noise_scheduler = kwargs.get('noise_scheduler', 'ddpm')
|
||||
self.content_or_style: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
|
||||
self.steps: int = kwargs.get('steps', 1000)
|
||||
self.lr = kwargs.get('lr', 1e-6)
|
||||
self.unet_lr = kwargs.get('unet_lr', self.lr)
|
||||
self.text_encoder_lr = kwargs.get('text_encoder_lr', self.lr)
|
||||
self.refiner_lr = kwargs.get('refiner_lr', self.lr)
|
||||
self.embedding_lr = kwargs.get('embedding_lr', self.lr)
|
||||
self.adapter_lr = kwargs.get('adapter_lr', self.lr)
|
||||
self.optimizer = kwargs.get('optimizer', 'adamw')
|
||||
self.optimizer_params = kwargs.get('optimizer_params', {})
|
||||
self.lr_scheduler = kwargs.get('lr_scheduler', 'constant')
|
||||
self.max_denoising_steps: int = kwargs.get('max_denoising_steps', 50)
|
||||
self.lr_scheduler_params = kwargs.get('lr_scheduler_params', {})
|
||||
self.min_denoising_steps: int = kwargs.get('min_denoising_steps', 0)
|
||||
self.max_denoising_steps: int = kwargs.get('max_denoising_steps', 1000)
|
||||
self.batch_size: int = kwargs.get('batch_size', 1)
|
||||
self.dtype: str = kwargs.get('dtype', 'fp32')
|
||||
self.xformers = kwargs.get('xformers', False)
|
||||
self.sdp = kwargs.get('sdp', False)
|
||||
self.train_unet = kwargs.get('train_unet', True)
|
||||
self.train_text_encoder = kwargs.get('train_text_encoder', True)
|
||||
self.train_refiner = kwargs.get('train_refiner', True)
|
||||
self.min_snr_gamma = kwargs.get('min_snr_gamma', None)
|
||||
self.snr_gamma = kwargs.get('snr_gamma', None)
|
||||
# trains a gamma, offset, and scale to adjust loss to adapt to timestep differentials
|
||||
# this should balance the learning rate across all timesteps over time
|
||||
self.learnable_snr_gos = kwargs.get('learnable_snr_gos', False)
|
||||
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)
|
||||
self.weight_jitter = kwargs.get('weight_jitter', 0.0)
|
||||
self.merge_network_on_save = kwargs.get('merge_network_on_save', False)
|
||||
self.max_grad_norm = kwargs.get('max_grad_norm', 1.0)
|
||||
self.start_step = kwargs.get('start_step', None)
|
||||
self.free_u = kwargs.get('free_u', False)
|
||||
self.adapter_assist_name_or_path: Optional[str] = kwargs.get('adapter_assist_name_or_path', None)
|
||||
self.noise_multiplier = kwargs.get('noise_multiplier', 1.0)
|
||||
self.img_multiplier = kwargs.get('img_multiplier', 1.0)
|
||||
self.latent_multiplier = kwargs.get('latent_multiplier', 1.0)
|
||||
self.negative_prompt = kwargs.get('negative_prompt', None)
|
||||
# multiplier applied to loos on regularization images
|
||||
self.reg_weight = kwargs.get('reg_weight', 1.0)
|
||||
|
||||
# dropout that happens before encoding. It functions independently per text encoder
|
||||
self.prompt_dropout_prob = kwargs.get('prompt_dropout_prob', 0.0)
|
||||
|
||||
# match the norm of the noise before computing loss. This will help the model maintain its
|
||||
# current understandin of the brightness of images.
|
||||
|
||||
self.match_noise_norm = kwargs.get('match_noise_norm', False)
|
||||
|
||||
# set to -1 to accumulate gradients for entire epoch
|
||||
# warning, only do this with a small dataset or you will run out of memory
|
||||
self.gradient_accumulation_steps = kwargs.get('gradient_accumulation_steps', 1)
|
||||
|
||||
# short long captions will double your batch size. This only works when a dataset is
|
||||
# prepared with a json caption file that has both short and long captions in it. It will
|
||||
# Double up every image and run it through with both short and long captions. The idea
|
||||
# is that the network will learn how to generate good images with both short and long captions
|
||||
self.short_and_long_captions = kwargs.get('short_and_long_captions', False)
|
||||
# if above is NOT true, this will make it so the long caption foes to te2 and the short caption goes to te1 for sdxl only
|
||||
self.short_and_long_captions_encoder_split = kwargs.get('short_and_long_captions_encoder_split', False)
|
||||
|
||||
# basically gradient accumulation but we run just 1 item through the network
|
||||
# and accumulate gradients. This can be used as basic gradient accumulation but is very helpful
|
||||
# for training tricks that increase batch size but need a single gradient step
|
||||
self.single_item_batching = kwargs.get('single_item_batching', False)
|
||||
|
||||
match_adapter_assist = kwargs.get('match_adapter_assist', False)
|
||||
self.match_adapter_chance = kwargs.get('match_adapter_chance', 0.0)
|
||||
self.loss_target: LossTarget = kwargs.get('loss_target',
|
||||
'noise') # noise, source, unaugmented, differential_noise
|
||||
|
||||
# When a mask is passed in a dataset, and this is true,
|
||||
# we will predict noise without a the LoRa network and use the prediction as a target for
|
||||
# unmasked reign. It is unmasked regularization basically
|
||||
self.inverted_mask_prior = kwargs.get('inverted_mask_prior', False)
|
||||
self.inverted_mask_prior_multiplier = kwargs.get('inverted_mask_prior_multiplier', 0.5)
|
||||
|
||||
# legacy
|
||||
if match_adapter_assist and self.match_adapter_chance == 0.0:
|
||||
self.match_adapter_chance = 1.0
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
@@ -68,11 +244,45 @@ class ModelConfig:
|
||||
self.name_or_path: str = kwargs.get('name_or_path', None)
|
||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||
self.is_ssd: bool = kwargs.get('is_ssd', False)
|
||||
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.vae_path = kwargs.get('vae_path', None)
|
||||
self.refiner_name_or_path = kwargs.get('refiner_name_or_path', None)
|
||||
self._original_refiner_name_or_path = self.refiner_name_or_path
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5)
|
||||
|
||||
# only for SDXL models for now
|
||||
self.use_text_encoder_1: bool = kwargs.get('use_text_encoder_1', True)
|
||||
self.use_text_encoder_2: bool = kwargs.get('use_text_encoder_2', True)
|
||||
|
||||
self.experimental_xl: bool = kwargs.get('experimental_xl', False)
|
||||
|
||||
if self.name_or_path is None:
|
||||
raise ValueError('name_or_path must be specified')
|
||||
|
||||
if self.is_ssd:
|
||||
# sed sdxl as true since it is mostly the same architecture
|
||||
self.is_xl = True
|
||||
|
||||
|
||||
class ReferenceDatasetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
# can pass with a side by side pait or a folder with pos and neg folder
|
||||
self.pair_folder: str = kwargs.get('pair_folder', None)
|
||||
self.pos_folder: str = kwargs.get('pos_folder', None)
|
||||
self.neg_folder: str = kwargs.get('neg_folder', None)
|
||||
|
||||
self.network_weight: float = float(kwargs.get('network_weight', 1.0))
|
||||
self.pos_weight: float = float(kwargs.get('pos_weight', self.network_weight))
|
||||
self.neg_weight: float = float(kwargs.get('neg_weight', self.network_weight))
|
||||
# make sure they are all absolute values no negatives
|
||||
self.pos_weight = abs(self.pos_weight)
|
||||
self.neg_weight = abs(self.neg_weight)
|
||||
|
||||
self.target_class: str = kwargs.get('target_class', '')
|
||||
self.size: int = kwargs.get('size', 512)
|
||||
|
||||
|
||||
class SliderTargetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -81,6 +291,15 @@ class SliderTargetConfig:
|
||||
self.negative: str = kwargs.get('negative', '')
|
||||
self.multiplier: float = kwargs.get('multiplier', 1.0)
|
||||
self.weight: float = kwargs.get('weight', 1.0)
|
||||
self.shuffle: bool = kwargs.get('shuffle', False)
|
||||
|
||||
|
||||
class GuidanceConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.target_class: str = kwargs.get('target_class', '')
|
||||
self.guidance_scale: float = kwargs.get('guidance_scale', 1.0)
|
||||
self.positive_prompt: str = kwargs.get('positive_prompt', '')
|
||||
self.negative_prompt: str = kwargs.get('negative_prompt', '')
|
||||
|
||||
|
||||
class SliderConfigAnchors:
|
||||
@@ -93,9 +312,332 @@ class SliderConfigAnchors:
|
||||
class SliderConfig:
|
||||
def __init__(self, **kwargs):
|
||||
targets = kwargs.get('targets', [])
|
||||
targets = [SliderTargetConfig(**target) for target in targets]
|
||||
self.targets: List[SliderTargetConfig] = targets
|
||||
anchors = kwargs.get('anchors', [])
|
||||
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)
|
||||
self.use_adapter: bool = kwargs.get('use_adapter', None) # depth
|
||||
self.adapter_img_dir = kwargs.get('adapter_img_dir', None)
|
||||
self.low_ram = kwargs.get('low_ram', False)
|
||||
|
||||
# expand targets if shuffling
|
||||
from toolkit.prompt_utils import get_slider_target_permutations
|
||||
self.targets: List[SliderTargetConfig] = []
|
||||
targets = [SliderTargetConfig(**target) for target in targets]
|
||||
# do permutations if shuffle is true
|
||||
print(f"Building slider targets")
|
||||
for target in targets:
|
||||
if target.shuffle:
|
||||
target_permutations = get_slider_target_permutations(target, max_permutations=8)
|
||||
self.targets = self.targets + target_permutations
|
||||
else:
|
||||
self.targets.append(target)
|
||||
print(f"Built {len(self.targets)} slider targets (with permutations)")
|
||||
|
||||
|
||||
class DatasetConfig:
|
||||
"""
|
||||
Dataset config for sd-datasets
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.type = kwargs.get('type', 'image') # sd, slider, reference
|
||||
# will be legacy
|
||||
self.folder_path: str = kwargs.get('folder_path', None)
|
||||
# can be json or folder path
|
||||
self.dataset_path: str = kwargs.get('dataset_path', None)
|
||||
|
||||
self.default_caption: str = kwargs.get('default_caption', None)
|
||||
self.random_triggers: List[str] = kwargs.get('random_triggers', [])
|
||||
self.caption_ext: str = kwargs.get('caption_ext', None)
|
||||
self.random_scale: bool = kwargs.get('random_scale', False)
|
||||
self.random_crop: bool = kwargs.get('random_crop', False)
|
||||
self.resolution: int = kwargs.get('resolution', 512)
|
||||
self.scale: float = kwargs.get('scale', 1.0)
|
||||
self.buckets: bool = kwargs.get('buckets', False)
|
||||
self.bucket_tolerance: int = kwargs.get('bucket_tolerance', 64)
|
||||
self.is_reg: bool = kwargs.get('is_reg', False)
|
||||
self.network_weight: float = float(kwargs.get('network_weight', 1.0))
|
||||
self.token_dropout_rate: float = float(kwargs.get('token_dropout_rate', 0.0))
|
||||
self.shuffle_tokens: bool = kwargs.get('shuffle_tokens', False)
|
||||
self.caption_dropout_rate: float = float(kwargs.get('caption_dropout_rate', 0.0))
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_y', False)
|
||||
self.augments: List[str] = kwargs.get('augments', [])
|
||||
self.control_path: str = kwargs.get('control_path', None) # depth maps, etc
|
||||
self.alpha_mask: bool = kwargs.get('alpha_mask', False) # if true, will use alpha channel as mask
|
||||
self.mask_path: str = kwargs.get('mask_path',
|
||||
None) # focus mask (black and white. White has higher loss than black)
|
||||
self.unconditional_path: str = kwargs.get('unconditional_path', None) # path where matching unconditional images are located
|
||||
self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask
|
||||
self.mask_min_value: float = kwargs.get('mask_min_value', 0.01) # min value for . 0 - 1
|
||||
self.poi: Union[str, None] = kwargs.get('poi',
|
||||
None) # if one is set and in json data, will be used as auto crop scale point of interes
|
||||
self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset
|
||||
# cache latents will store them in memory
|
||||
self.cache_latents: bool = kwargs.get('cache_latents', False)
|
||||
# cache latents to disk will store them on disk. If both are true, it will save to disk, but keep in memory
|
||||
self.cache_latents_to_disk: bool = kwargs.get('cache_latents_to_disk', False)
|
||||
|
||||
# https://albumentations.ai/docs/api_reference/augmentations/transforms
|
||||
# augmentations are returned as a separate image and cannot currently be cached
|
||||
self.augmentations: List[dict] = kwargs.get('augmentations', None)
|
||||
self.shuffle_augmentations: bool = kwargs.get('shuffle_augmentations', False)
|
||||
|
||||
has_augmentations = self.augmentations is not None and len(self.augmentations) > 0
|
||||
|
||||
if (len(self.augments) > 0 or has_augmentations) and (self.cache_latents or self.cache_latents_to_disk):
|
||||
print(f"WARNING: Augments are not supported with caching latents. Setting cache_latents to False")
|
||||
self.cache_latents = False
|
||||
self.cache_latents_to_disk = False
|
||||
|
||||
# legacy compatability
|
||||
legacy_caption_type = kwargs.get('caption_type', None)
|
||||
if legacy_caption_type:
|
||||
self.caption_ext = legacy_caption_type
|
||||
self.caption_type = self.caption_ext
|
||||
|
||||
|
||||
def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]:
|
||||
"""
|
||||
This just splits up the datasets by resolutions so you dont have to do it manually
|
||||
:param raw_config:
|
||||
:return:
|
||||
"""
|
||||
# split up datasets by resolutions
|
||||
new_config = []
|
||||
for dataset in raw_config:
|
||||
resolution = dataset.get('resolution', 512)
|
||||
if isinstance(resolution, list):
|
||||
resolution_list = resolution
|
||||
else:
|
||||
resolution_list = [resolution]
|
||||
for res in resolution_list:
|
||||
dataset_copy = dataset.copy()
|
||||
dataset_copy['resolution'] = res
|
||||
new_config.append(dataset_copy)
|
||||
return new_config
|
||||
|
||||
|
||||
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 = ImgExt, # 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
|
||||
adapter_image_path: str = None, # path to adapter image
|
||||
adapter_conditioning_scale: float = 1.0, # scale for adapter conditioning
|
||||
latents: Union[torch.Tensor | None] = None, # input latent to start with,
|
||||
extra_kwargs: dict = None, # extra data to save with prompt file
|
||||
refiner_start_at: float = 0.5, # start at this percentage of a step. 0.0 to 1.0 . 1.0 is the end
|
||||
):
|
||||
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.latents: Union[torch.Tensor | None] = latents
|
||||
|
||||
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)
|
||||
self.adapter_image_path: str = adapter_image_path
|
||||
self.adapter_conditioning_scale: float = adapter_conditioning_scale
|
||||
self.extra_kwargs = extra_kwargs if extra_kwargs is not None else {}
|
||||
self.refiner_start_at = refiner_start_at
|
||||
|
||||
# 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)
|
||||
ext = self.output_ext
|
||||
# if it does not start with a dot add one
|
||||
if ext[0] != '.':
|
||||
ext = '.' + ext
|
||||
filename += 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)
|
||||
elif flag == 'a':
|
||||
self.adapter_conditioning_scale = float(content)
|
||||
elif flag == 'ref':
|
||||
self.refiner_start_at = float(content)
|
||||
|
||||
def post_process_embeddings(
|
||||
self,
|
||||
conditional_prompt_embeds: PromptEmbeds,
|
||||
unconditional_prompt_embeds: Optional[PromptEmbeds] = None,
|
||||
):
|
||||
# this is called after prompt embeds are encoded. We can override them in the future here
|
||||
pass
|
||||
|
||||
93
toolkit/cuda_malloc.py
Normal file
93
toolkit/cuda_malloc.py
Normal file
@@ -0,0 +1,93 @@
|
||||
# ref comfy ui
|
||||
import os
|
||||
import importlib.util
|
||||
|
||||
|
||||
# Can't use pytorch to get the GPU names because the cuda malloc has to be set before the first import.
|
||||
def get_gpu_names():
|
||||
if os.name == 'nt':
|
||||
import ctypes
|
||||
|
||||
# Define necessary C structures and types
|
||||
class DISPLAY_DEVICEA(ctypes.Structure):
|
||||
_fields_ = [
|
||||
('cb', ctypes.c_ulong),
|
||||
('DeviceName', ctypes.c_char * 32),
|
||||
('DeviceString', ctypes.c_char * 128),
|
||||
('StateFlags', ctypes.c_ulong),
|
||||
('DeviceID', ctypes.c_char * 128),
|
||||
('DeviceKey', ctypes.c_char * 128)
|
||||
]
|
||||
|
||||
# Load user32.dll
|
||||
user32 = ctypes.windll.user32
|
||||
|
||||
# Call EnumDisplayDevicesA
|
||||
def enum_display_devices():
|
||||
device_info = DISPLAY_DEVICEA()
|
||||
device_info.cb = ctypes.sizeof(device_info)
|
||||
device_index = 0
|
||||
gpu_names = set()
|
||||
|
||||
while user32.EnumDisplayDevicesA(None, device_index, ctypes.byref(device_info), 0):
|
||||
device_index += 1
|
||||
gpu_names.add(device_info.DeviceString.decode('utf-8'))
|
||||
return gpu_names
|
||||
|
||||
return enum_display_devices()
|
||||
else:
|
||||
return set()
|
||||
|
||||
|
||||
blacklist = {"GeForce GTX TITAN X", "GeForce GTX 980", "GeForce GTX 970", "GeForce GTX 960", "GeForce GTX 950",
|
||||
"GeForce 945M",
|
||||
"GeForce 940M", "GeForce 930M", "GeForce 920M", "GeForce 910M", "GeForce GTX 750", "GeForce GTX 745",
|
||||
"Quadro K620",
|
||||
"Quadro K1200", "Quadro K2200", "Quadro M500", "Quadro M520", "Quadro M600", "Quadro M620", "Quadro M1000",
|
||||
"Quadro M1200", "Quadro M2000", "Quadro M2200", "Quadro M3000", "Quadro M4000", "Quadro M5000",
|
||||
"Quadro M5500", "Quadro M6000",
|
||||
"GeForce MX110", "GeForce MX130", "GeForce 830M", "GeForce 840M", "GeForce GTX 850M", "GeForce GTX 860M",
|
||||
"GeForce GTX 1650", "GeForce GTX 1630"
|
||||
}
|
||||
|
||||
|
||||
def cuda_malloc_supported():
|
||||
try:
|
||||
names = get_gpu_names()
|
||||
except:
|
||||
names = set()
|
||||
for x in names:
|
||||
if "NVIDIA" in x:
|
||||
for b in blacklist:
|
||||
if b in x:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
cuda_malloc = False
|
||||
|
||||
if not cuda_malloc:
|
||||
try:
|
||||
version = ""
|
||||
torch_spec = importlib.util.find_spec("torch")
|
||||
for folder in torch_spec.submodule_search_locations:
|
||||
ver_file = os.path.join(folder, "version.py")
|
||||
if os.path.isfile(ver_file):
|
||||
spec = importlib.util.spec_from_file_location("torch_version_import", ver_file)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
version = module.__version__
|
||||
if int(version[0]) >= 2: # enable by default for torch version 2.0 and up
|
||||
cuda_malloc = cuda_malloc_supported()
|
||||
except:
|
||||
pass
|
||||
|
||||
if cuda_malloc:
|
||||
env_var = os.environ.get('PYTORCH_CUDA_ALLOC_CONF', None)
|
||||
if env_var is None:
|
||||
env_var = "backend:cudaMallocAsync"
|
||||
else:
|
||||
env_var += ",backend:cudaMallocAsync"
|
||||
|
||||
os.environ['PYTORCH_CUDA_ALLOC_CONF'] = env_var
|
||||
print("CUDA Malloc Async Enabled")
|
||||
@@ -1,19 +1,42 @@
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import traceback
|
||||
from functools import lru_cache
|
||||
from typing import List, TYPE_CHECKING
|
||||
|
||||
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 torch.utils.data import Dataset, DataLoader, ConcatDataset
|
||||
from tqdm import tqdm
|
||||
import albumentations as A
|
||||
|
||||
from toolkit.buckets import get_bucket_for_image_size, BucketResolution
|
||||
from toolkit.config_modules import DatasetConfig, preprocess_dataset_raw_config
|
||||
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class ImageDataset(Dataset):
|
||||
class ImageDataset(Dataset, CaptionMixin):
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.name = self.get_config('name', 'dataset')
|
||||
self.path = self.get_config('path', required=True)
|
||||
self.scale = self.get_config('scale', 1)
|
||||
self.random_scale = self.get_config('random_scale', False)
|
||||
self.include_prompt = self.get_config('include_prompt', False)
|
||||
self.default_prompt = self.get_config('default_prompt', '')
|
||||
if self.include_prompt:
|
||||
self.caption_type = self.get_config('caption_ext', 'txt')
|
||||
else:
|
||||
self.caption_type = None
|
||||
# we always random crop if random scale is enabled
|
||||
self.random_crop = self.random_scale if self.random_scale else self.get_config('random_crop', False)
|
||||
|
||||
@@ -32,13 +55,15 @@ class ImageDataset(Dataset):
|
||||
else:
|
||||
bad_count += 1
|
||||
|
||||
self.file_list = new_file_list
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
print(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.path}"
|
||||
|
||||
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,11 +90,14 @@ 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))
|
||||
img = img.resize((scale_size, scale_size), Image.BICUBIC)
|
||||
scaler = scale_size / min_img_size
|
||||
scale_width = int((img.width + 5) * scaler)
|
||||
scale_height = int((img.height + 5) * scaler)
|
||||
img = img.resize((scale_width, scale_height), Image.BICUBIC)
|
||||
img = transforms.RandomCrop(self.resolution)(img)
|
||||
else:
|
||||
img = transforms.CenterCrop(min_img_size)(img)
|
||||
@@ -77,4 +105,454 @@ class ImageDataset(Dataset):
|
||||
|
||||
img = self.transform(img)
|
||||
|
||||
return img
|
||||
if self.include_prompt:
|
||||
prompt = self.get_caption_item(index)
|
||||
return img, prompt
|
||||
else:
|
||||
return img
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
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', None)
|
||||
self.pos_folder = self.get_config('pos_folder', None)
|
||||
self.neg_folder = self.get_config('neg_folder', None)
|
||||
|
||||
self.default_prompt = self.get_config('default_prompt', '')
|
||||
self.network_weight = self.get_config('network_weight', 1.0)
|
||||
self.pos_weight = self.get_config('pos_weight', self.network_weight)
|
||||
self.neg_weight = self.get_config('neg_weight', self.network_weight)
|
||||
|
||||
supported_exts = ('.jpg', '.jpeg', '.png', '.webp', '.JPEG', '.JPG', '.PNG', '.WEBP')
|
||||
|
||||
if self.pos_folder is not None and self.neg_folder is not None:
|
||||
# find matching files
|
||||
self.pos_file_list = [os.path.join(self.pos_folder, file) for file in os.listdir(self.pos_folder) if
|
||||
file.lower().endswith(supported_exts)]
|
||||
self.neg_file_list = [os.path.join(self.neg_folder, file) for file in os.listdir(self.neg_folder) if
|
||||
file.lower().endswith(supported_exts)]
|
||||
|
||||
matched_files = []
|
||||
for pos_file in self.pos_file_list:
|
||||
pos_file_no_ext = os.path.splitext(pos_file)[0]
|
||||
for neg_file in self.neg_file_list:
|
||||
neg_file_no_ext = os.path.splitext(neg_file)[0]
|
||||
if os.path.basename(pos_file_no_ext) == os.path.basename(neg_file_no_ext):
|
||||
matched_files.append((neg_file, pos_file))
|
||||
break
|
||||
|
||||
# remove duplicates
|
||||
matched_files = [t for t in (set(tuple(i) for i in matched_files))]
|
||||
|
||||
self.file_list = matched_files
|
||||
print(f" - Found {len(self.file_list)} matching pairs")
|
||||
else:
|
||||
self.file_list = [os.path.join(self.path, file) for file in os.listdir(self.path) if
|
||||
file.lower().endswith(supported_exts)]
|
||||
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 get_all_prompts(self):
|
||||
prompts = []
|
||||
for index in range(len(self.file_list)):
|
||||
prompts.append(self.get_prompt_item(index))
|
||||
|
||||
# remove duplicates
|
||||
prompts = list(set(prompts))
|
||||
return prompts
|
||||
|
||||
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 get_prompt_item(self, index):
|
||||
img_path_or_tuple = self.file_list[index]
|
||||
if isinstance(img_path_or_tuple, tuple):
|
||||
# check if either has a prompt file
|
||||
path_no_ext = os.path.splitext(img_path_or_tuple[0])[0]
|
||||
prompt_path = path_no_ext + '.txt'
|
||||
if not os.path.exists(prompt_path):
|
||||
path_no_ext = os.path.splitext(img_path_or_tuple[1])[0]
|
||||
prompt_path = path_no_ext + '.txt'
|
||||
else:
|
||||
img_path = img_path_or_tuple
|
||||
# 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
|
||||
return prompt
|
||||
|
||||
def __getitem__(self, index):
|
||||
img_path_or_tuple = self.file_list[index]
|
||||
if isinstance(img_path_or_tuple, tuple):
|
||||
# load both images
|
||||
img_path = img_path_or_tuple[0]
|
||||
img1 = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
img_path = img_path_or_tuple[1]
|
||||
img2 = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
|
||||
# always use # 2 (pos)
|
||||
bucket_resolution = get_bucket_for_image_size(
|
||||
width=img2.width,
|
||||
height=img2.height,
|
||||
resolution=self.size,
|
||||
# divisibility=self.
|
||||
)
|
||||
|
||||
# images will be same base dimension, but may be trimmed. We need to shrink and then central crop
|
||||
if bucket_resolution['width'] > bucket_resolution['height']:
|
||||
img1_scale_to_height = bucket_resolution["height"]
|
||||
img1_scale_to_width = int(img1.width * (bucket_resolution["height"] / img1.height))
|
||||
img2_scale_to_height = bucket_resolution["height"]
|
||||
img2_scale_to_width = int(img2.width * (bucket_resolution["height"] / img2.height))
|
||||
else:
|
||||
img1_scale_to_width = bucket_resolution["width"]
|
||||
img1_scale_to_height = int(img1.height * (bucket_resolution["width"] / img1.width))
|
||||
img2_scale_to_width = bucket_resolution["width"]
|
||||
img2_scale_to_height = int(img2.height * (bucket_resolution["width"] / img2.width))
|
||||
|
||||
img1_crop_height = bucket_resolution["height"]
|
||||
img1_crop_width = bucket_resolution["width"]
|
||||
img2_crop_height = bucket_resolution["height"]
|
||||
img2_crop_width = bucket_resolution["width"]
|
||||
|
||||
# scale then center crop images
|
||||
img1 = img1.resize((img1_scale_to_width, img1_scale_to_height), Image.BICUBIC)
|
||||
img1 = transforms.CenterCrop((img1_crop_height, img1_crop_width))(img1)
|
||||
img2 = img2.resize((img2_scale_to_width, img2_scale_to_height), Image.BICUBIC)
|
||||
img2 = transforms.CenterCrop((img2_crop_height, img2_crop_width))(img2)
|
||||
|
||||
# combine them side by side
|
||||
img = Image.new('RGB', (img1.width + img2.width, max(img1.height, img2.height)))
|
||||
img.paste(img1, (0, 0))
|
||||
img.paste(img2, (img1.width, 0))
|
||||
else:
|
||||
img_path = img_path_or_tuple
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
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)
|
||||
|
||||
prompt = self.get_prompt_item(index)
|
||||
img = self.transform(img)
|
||||
|
||||
return img, prompt, (self.neg_weight, self.pos_weight)
|
||||
|
||||
|
||||
class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dataset_config: 'DatasetConfig',
|
||||
batch_size=1,
|
||||
sd: 'StableDiffusion' = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dataset_config = dataset_config
|
||||
folder_path = dataset_config.folder_path
|
||||
self.dataset_path = dataset_config.dataset_path
|
||||
if self.dataset_path is None:
|
||||
self.dataset_path = folder_path
|
||||
|
||||
self.is_caching_latents = dataset_config.cache_latents or dataset_config.cache_latents_to_disk
|
||||
self.is_caching_latents_to_memory = dataset_config.cache_latents
|
||||
self.is_caching_latents_to_disk = dataset_config.cache_latents_to_disk
|
||||
self.epoch_num = 0
|
||||
|
||||
self.sd = sd
|
||||
|
||||
if self.sd is None and self.is_caching_latents:
|
||||
raise ValueError(f"sd is required for caching latents")
|
||||
|
||||
self.caption_type = dataset_config.caption_ext
|
||||
self.default_caption = dataset_config.default_caption
|
||||
self.random_scale = dataset_config.random_scale
|
||||
self.scale = dataset_config.scale
|
||||
self.batch_size = batch_size
|
||||
# we always random crop if random scale is enabled
|
||||
self.random_crop = self.random_scale if self.random_scale else dataset_config.random_crop
|
||||
self.resolution = dataset_config.resolution
|
||||
self.caption_dict = None
|
||||
self.file_list: List['FileItemDTO'] = []
|
||||
|
||||
# check if dataset_path is a folder or json
|
||||
if os.path.isdir(self.dataset_path):
|
||||
file_list = [
|
||||
os.path.join(self.dataset_path, file) for file in os.listdir(self.dataset_path) if
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))
|
||||
]
|
||||
else:
|
||||
# assume json
|
||||
with open(self.dataset_path, 'r') as f:
|
||||
self.caption_dict = json.load(f)
|
||||
# keys are file paths
|
||||
file_list = list(self.caption_dict.keys())
|
||||
|
||||
if self.dataset_config.num_repeats > 1:
|
||||
# repeat the list
|
||||
file_list = file_list * self.dataset_config.num_repeats
|
||||
|
||||
# this might take a while
|
||||
print(f" - Preprocessing image dimensions")
|
||||
bad_count = 0
|
||||
for file in tqdm(file_list):
|
||||
try:
|
||||
file_item = FileItemDTO(
|
||||
path=file,
|
||||
dataset_config=dataset_config
|
||||
)
|
||||
self.file_list.append(file_item)
|
||||
except Exception as e:
|
||||
print(traceback.format_exc())
|
||||
print(f"Error processing image: {file}")
|
||||
print(e)
|
||||
bad_count += 1
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
# print(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
|
||||
|
||||
# handle x axis flips
|
||||
if self.dataset_config.flip_x:
|
||||
print(" - adding x axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the x axis
|
||||
new_file_item = copy.deepcopy(file_item)
|
||||
new_file_item.flip_x = True
|
||||
self.file_list.append(new_file_item)
|
||||
|
||||
# handle y axis flips
|
||||
if self.dataset_config.flip_y:
|
||||
print(" - adding y axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the y axis
|
||||
new_file_item = copy.deepcopy(file_item)
|
||||
new_file_item.flip_y = True
|
||||
self.file_list.append(new_file_item)
|
||||
|
||||
if self.dataset_config.flip_x or self.dataset_config.flip_y:
|
||||
print(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
self.setup_epoch()
|
||||
|
||||
def setup_epoch(self):
|
||||
if self.epoch_num == 0:
|
||||
# initial setup
|
||||
# do not call for now
|
||||
if self.dataset_config.buckets:
|
||||
# setup buckets
|
||||
self.setup_buckets()
|
||||
if self.is_caching_latents:
|
||||
self.cache_latents_all_latents()
|
||||
else:
|
||||
if self.dataset_config.poi is not None:
|
||||
# handle cropping to a specific point of interest
|
||||
# setup buckets every epoch
|
||||
self.setup_buckets(quiet=True)
|
||||
self.epoch_num += 1
|
||||
|
||||
def __len__(self):
|
||||
if self.dataset_config.buckets:
|
||||
return len(self.batch_indices)
|
||||
return len(self.file_list)
|
||||
|
||||
def _get_single_item(self, index) -> 'FileItemDTO':
|
||||
file_item = copy.deepcopy(self.file_list[index])
|
||||
file_item.load_and_process_image(self.transform)
|
||||
file_item.load_caption(self.caption_dict)
|
||||
return file_item
|
||||
|
||||
def __getitem__(self, item):
|
||||
if self.dataset_config.buckets:
|
||||
# for buckets we collate ourselves for now
|
||||
# todo allow a scheduler to dynamically make buckets
|
||||
# we collate ourselves
|
||||
if len(self.batch_indices) - 1 < item:
|
||||
# tried everything to solve this. No way to reset length when redoing things. Pick another index
|
||||
item = random.randint(0, len(self.batch_indices) - 1)
|
||||
idx_list = self.batch_indices[item]
|
||||
return [self._get_single_item(idx) for idx in idx_list]
|
||||
else:
|
||||
# Dataloader is batching
|
||||
return self._get_single_item(item)
|
||||
|
||||
|
||||
def get_dataloader_from_datasets(
|
||||
dataset_options,
|
||||
batch_size=1,
|
||||
sd: 'StableDiffusion' = None,
|
||||
) -> DataLoader:
|
||||
if dataset_options is None or len(dataset_options) == 0:
|
||||
return None
|
||||
|
||||
datasets = []
|
||||
has_buckets = False
|
||||
is_caching_latents = False
|
||||
|
||||
dataset_config_list = []
|
||||
# preprocess them all
|
||||
for dataset_option in dataset_options:
|
||||
if isinstance(dataset_option, DatasetConfig):
|
||||
dataset_config_list.append(dataset_option)
|
||||
else:
|
||||
# preprocess raw data
|
||||
split_configs = preprocess_dataset_raw_config([dataset_option])
|
||||
for x in split_configs:
|
||||
dataset_config_list.append(DatasetConfig(**x))
|
||||
|
||||
for config in dataset_config_list:
|
||||
|
||||
if config.type == 'image':
|
||||
dataset = AiToolkitDataset(config, batch_size=batch_size, sd=sd)
|
||||
datasets.append(dataset)
|
||||
if config.buckets:
|
||||
has_buckets = True
|
||||
if config.cache_latents or config.cache_latents_to_disk:
|
||||
is_caching_latents = True
|
||||
else:
|
||||
raise ValueError(f"invalid dataset type: {config.type}")
|
||||
|
||||
concatenated_dataset = ConcatDataset(datasets)
|
||||
|
||||
# todo build scheduler that can get buckets from all datasets that match
|
||||
# todo and evenly distribute reg images
|
||||
|
||||
def dto_collation(batch: List['FileItemDTO']):
|
||||
# create DTO batch
|
||||
batch = DataLoaderBatchDTO(
|
||||
file_items=batch
|
||||
)
|
||||
return batch
|
||||
|
||||
# check if is caching latents
|
||||
|
||||
|
||||
if has_buckets:
|
||||
# make sure they all have buckets
|
||||
for dataset in datasets:
|
||||
assert dataset.dataset_config.buckets, f"buckets not found on dataset {dataset.dataset_config.folder_path}, you either need all buckets or none"
|
||||
|
||||
data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=None, # we batch in the datasets for now
|
||||
drop_last=False,
|
||||
shuffle=True,
|
||||
collate_fn=dto_collation, # Use the custom collate function
|
||||
num_workers=4
|
||||
)
|
||||
else:
|
||||
data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=True,
|
||||
num_workers=4,
|
||||
collate_fn=dto_collation
|
||||
)
|
||||
return data_loader
|
||||
|
||||
|
||||
def trigger_dataloader_setup_epoch(dataloader: DataLoader):
|
||||
# hacky but needed because of different types of datasets and dataloaders
|
||||
dataloader.len = None
|
||||
if isinstance(dataloader.dataset, list):
|
||||
for dataset in dataloader.dataset:
|
||||
if hasattr(dataset, 'datasets'):
|
||||
for sub_dataset in dataset.datasets:
|
||||
if hasattr(sub_dataset, 'setup_epoch'):
|
||||
sub_dataset.setup_epoch()
|
||||
sub_dataset.len = None
|
||||
elif hasattr(dataset, 'setup_epoch'):
|
||||
dataset.setup_epoch()
|
||||
dataset.len = None
|
||||
elif hasattr(dataloader.dataset, 'setup_epoch'):
|
||||
dataloader.dataset.setup_epoch()
|
||||
dataloader.dataset.len = None
|
||||
elif hasattr(dataloader.dataset, 'datasets'):
|
||||
dataloader.dataset.len = None
|
||||
for sub_dataset in dataloader.dataset.datasets:
|
||||
if hasattr(sub_dataset, 'setup_epoch'):
|
||||
sub_dataset.setup_epoch()
|
||||
sub_dataset.len = None
|
||||
|
||||
202
toolkit/data_transfer_object/data_loader.py
Normal file
202
toolkit/data_transfer_object/data_loader.py
Normal file
@@ -0,0 +1,202 @@
|
||||
from typing import TYPE_CHECKING, List, Union
|
||||
import torch
|
||||
import random
|
||||
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
|
||||
from toolkit import image_utils
|
||||
from toolkit.dataloader_mixins import CaptionProcessingDTOMixin, ImageProcessingDTOMixin, LatentCachingFileItemDTOMixin, \
|
||||
ControlFileItemDTOMixin, ArgBreakMixin, PoiFileItemDTOMixin, MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, \
|
||||
UnconditionalFileItemDTOMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
|
||||
printed_messages = []
|
||||
|
||||
|
||||
def print_once(msg):
|
||||
global printed_messages
|
||||
if msg not in printed_messages:
|
||||
print(msg)
|
||||
printed_messages.append(msg)
|
||||
|
||||
|
||||
class FileItemDTO(
|
||||
LatentCachingFileItemDTOMixin,
|
||||
CaptionProcessingDTOMixin,
|
||||
ImageProcessingDTOMixin,
|
||||
ControlFileItemDTOMixin,
|
||||
MaskFileItemDTOMixin,
|
||||
AugmentationFileItemDTOMixin,
|
||||
UnconditionalFileItemDTOMixin,
|
||||
PoiFileItemDTOMixin,
|
||||
ArgBreakMixin,
|
||||
):
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.path = kwargs.get('path', None)
|
||||
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
# process width and height
|
||||
try:
|
||||
w, h = image_utils.get_image_size(self.path)
|
||||
except image_utils.UnknownImageFormat:
|
||||
print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
|
||||
f'This process is faster for png, jpeg')
|
||||
img = exif_transpose(Image.open(self.path))
|
||||
h, w = img.size
|
||||
self.width: int = w
|
||||
self.height: int = h
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
# self.caption_path: str = kwargs.get('caption_path', None)
|
||||
self.raw_caption: str = kwargs.get('raw_caption', None)
|
||||
# we scale first, then crop
|
||||
self.scale_to_width: int = kwargs.get('scale_to_width', int(self.width * self.dataset_config.scale))
|
||||
self.scale_to_height: int = kwargs.get('scale_to_height', int(self.height * self.dataset_config.scale))
|
||||
# crop values are from scaled size
|
||||
self.crop_x: int = kwargs.get('crop_x', 0)
|
||||
self.crop_y: int = kwargs.get('crop_y', 0)
|
||||
self.crop_width: int = kwargs.get('crop_width', self.scale_to_width)
|
||||
self.crop_height: int = kwargs.get('crop_height', self.scale_to_height)
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_x', False)
|
||||
self.augments: List[str] = self.dataset_config.augments
|
||||
|
||||
self.network_weight: float = self.dataset_config.network_weight
|
||||
self.is_reg = self.dataset_config.is_reg
|
||||
self.tensor: Union[torch.Tensor, None] = None
|
||||
|
||||
def cleanup(self):
|
||||
self.tensor = None
|
||||
self.cleanup_latent()
|
||||
self.cleanup_control()
|
||||
self.cleanup_mask()
|
||||
self.cleanup_unconditional()
|
||||
|
||||
|
||||
class DataLoaderBatchDTO:
|
||||
def __init__(self, **kwargs):
|
||||
try:
|
||||
self.file_items: List['FileItemDTO'] = kwargs.get('file_items', None)
|
||||
is_latents_cached = self.file_items[0].is_latent_cached
|
||||
self.tensor: Union[torch.Tensor, None] = None
|
||||
self.latents: Union[torch.Tensor, None] = None
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
self.mask_tensor: Union[torch.Tensor, None] = None
|
||||
self.unaugmented_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_latents: Union[torch.Tensor, None] = None
|
||||
self.sigmas: Union[torch.Tensor, None] = None # can be added elseware and passed along training code
|
||||
if not is_latents_cached:
|
||||
# only return a tensor if latents are not cached
|
||||
self.tensor: torch.Tensor = torch.cat([x.tensor.unsqueeze(0) for x in self.file_items])
|
||||
# if we have encoded latents, we concatenate them
|
||||
self.latents: Union[torch.Tensor, None] = None
|
||||
if is_latents_cached:
|
||||
self.latents = torch.cat([x.get_latent().unsqueeze(0) for x in self.file_items])
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
# if self.file_items[0].control_tensor is not None:
|
||||
# if any have a control tensor, we concatenate them
|
||||
if any([x.control_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_control_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.control_tensor is not None:
|
||||
base_control_tensor = x.control_tensor
|
||||
break
|
||||
control_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.control_tensor is None:
|
||||
control_tensors.append(torch.zeros_like(base_control_tensor))
|
||||
else:
|
||||
control_tensors.append(x.control_tensor)
|
||||
self.control_tensor = torch.cat([x.unsqueeze(0) for x in control_tensors])
|
||||
|
||||
if any([x.mask_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_mask_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.mask_tensor is not None:
|
||||
base_mask_tensor = x.mask_tensor
|
||||
break
|
||||
mask_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.mask_tensor is None:
|
||||
mask_tensors.append(torch.zeros_like(base_mask_tensor))
|
||||
else:
|
||||
mask_tensors.append(x.mask_tensor)
|
||||
self.mask_tensor = torch.cat([x.unsqueeze(0) for x in mask_tensors])
|
||||
|
||||
# add unaugmented tensors for ones with augments
|
||||
if any([x.unaugmented_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_unaugmented_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.unaugmented_tensor is not None:
|
||||
base_unaugmented_tensor = x.unaugmented_tensor
|
||||
break
|
||||
unaugmented_tensor = []
|
||||
for x in self.file_items:
|
||||
if x.unaugmented_tensor is None:
|
||||
unaugmented_tensor.append(torch.zeros_like(base_unaugmented_tensor))
|
||||
else:
|
||||
unaugmented_tensor.append(x.unaugmented_tensor)
|
||||
self.unaugmented_tensor = torch.cat([x.unsqueeze(0) for x in unaugmented_tensor])
|
||||
|
||||
# add unconditional tensors
|
||||
if any([x.unconditional_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_unconditional_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.unaugmented_tensor is not None:
|
||||
base_unconditional_tensor = x.unconditional_tensor
|
||||
break
|
||||
unconditional_tensor = []
|
||||
for x in self.file_items:
|
||||
if x.unconditional_tensor is None:
|
||||
unconditional_tensor.append(torch.zeros_like(base_unconditional_tensor))
|
||||
else:
|
||||
unconditional_tensor.append(x.unconditional_tensor)
|
||||
self.unconditional_tensor = torch.cat([x.unsqueeze(0) for x in unconditional_tensor])
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
|
||||
def get_is_reg_list(self):
|
||||
return [x.is_reg for x in self.file_items]
|
||||
|
||||
def get_network_weight_list(self):
|
||||
return [x.network_weight for x in self.file_items]
|
||||
|
||||
def get_caption_list(
|
||||
self,
|
||||
trigger=None,
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present
|
||||
) for x in self.file_items]
|
||||
|
||||
def get_caption_short_list(
|
||||
self,
|
||||
trigger=None,
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present,
|
||||
short_caption=True
|
||||
) for x in self.file_items]
|
||||
|
||||
def cleanup(self):
|
||||
del self.latents
|
||||
del self.tensor
|
||||
del self.control_tensor
|
||||
for file_item in self.file_items:
|
||||
file_item.cleanup()
|
||||
1042
toolkit/dataloader_mixins.py
Normal file
1042
toolkit/dataloader_mixins.py
Normal file
File diff suppressed because it is too large
Load Diff
283
toolkit/embedding.py
Normal file
283
toolkit/embedding.py
Normal file
@@ -0,0 +1,283 @@
|
||||
import json
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
|
||||
import safetensors
|
||||
import torch
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.config_modules import EmbeddingConfig
|
||||
|
||||
|
||||
# this is a frankenstein mix of automatic1111 and my own code
|
||||
|
||||
class Embedding:
|
||||
def __init__(
|
||||
self,
|
||||
sd: 'StableDiffusion',
|
||||
embed_config: 'EmbeddingConfig',
|
||||
state_dict: OrderedDict = None,
|
||||
):
|
||||
self.name = embed_config.trigger
|
||||
self.sd = sd
|
||||
self.trigger = embed_config.trigger
|
||||
self.embed_config = embed_config
|
||||
self.step = 0
|
||||
# setup our embedding
|
||||
# Add the placeholder token in tokenizer
|
||||
placeholder_tokens = [self.embed_config.trigger]
|
||||
|
||||
# add dummy tokens for multi-vector
|
||||
additional_tokens = []
|
||||
for i in range(1, self.embed_config.tokens):
|
||||
additional_tokens.append(f"{self.embed_config.trigger}_{i}")
|
||||
placeholder_tokens += additional_tokens
|
||||
|
||||
# handle dual tokenizer
|
||||
self.tokenizer_list = self.sd.tokenizer if isinstance(self.sd.tokenizer, list) else [self.sd.tokenizer]
|
||||
self.text_encoder_list = self.sd.text_encoder if isinstance(self.sd.text_encoder, list) else [
|
||||
self.sd.text_encoder]
|
||||
|
||||
self.placeholder_token_ids = []
|
||||
self.embedding_tokens = []
|
||||
|
||||
print(f"Adding {placeholder_tokens} tokens to tokenizer")
|
||||
print(f"Adding {self.embed_config.tokens} tokens to tokenizer")
|
||||
|
||||
for text_encoder, tokenizer in zip(self.text_encoder_list, self.tokenizer_list):
|
||||
num_added_tokens = tokenizer.add_tokens(placeholder_tokens)
|
||||
if num_added_tokens != self.embed_config.tokens:
|
||||
raise ValueError(
|
||||
f"The tokenizer already contains the token {self.embed_config.trigger}. Please pass a different"
|
||||
f" `placeholder_token` that is not already in the tokenizer. Only added {num_added_tokens}"
|
||||
)
|
||||
|
||||
# Convert the initializer_token, placeholder_token to ids
|
||||
init_token_ids = tokenizer.encode(self.embed_config.init_words, add_special_tokens=False)
|
||||
# if length of token ids is more than number of orm embedding tokens fill with *
|
||||
if len(init_token_ids) > self.embed_config.tokens:
|
||||
init_token_ids = init_token_ids[:self.embed_config.tokens]
|
||||
elif len(init_token_ids) < self.embed_config.tokens:
|
||||
pad_token_id = tokenizer.encode(["*"], add_special_tokens=False)
|
||||
init_token_ids += pad_token_id * (self.embed_config.tokens - len(init_token_ids))
|
||||
|
||||
placeholder_token_ids = tokenizer.encode(placeholder_tokens, add_special_tokens=False)
|
||||
self.placeholder_token_ids.append(placeholder_token_ids)
|
||||
|
||||
# Resize the token embeddings as we are adding new special tokens to the tokenizer
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
# Initialise the newly added placeholder token with the embeddings of the initializer token
|
||||
token_embeds = text_encoder.get_input_embeddings().weight.data
|
||||
with torch.no_grad():
|
||||
for initializer_token_id, token_id in zip(init_token_ids, placeholder_token_ids):
|
||||
token_embeds[token_id] = token_embeds[initializer_token_id].clone()
|
||||
|
||||
# replace "[name] with this. on training. This is automatically generated in pipeline on inference
|
||||
self.embedding_tokens.append(" ".join(tokenizer.convert_ids_to_tokens(placeholder_token_ids)))
|
||||
|
||||
# backup text encoder embeddings
|
||||
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
|
||||
|
||||
def restore_embeddings(self):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[
|
||||
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
with torch.no_grad():
|
||||
text_encoder.get_input_embeddings().weight[
|
||||
index_no_updates
|
||||
] = orig_embeds[index_no_updates]
|
||||
|
||||
def get_trainable_params(self):
|
||||
params = []
|
||||
for text_encoder in self.text_encoder_list:
|
||||
params += text_encoder.get_input_embeddings().parameters()
|
||||
return params
|
||||
|
||||
def _get_vec(self, text_encoder_idx=0):
|
||||
# should we get params instead
|
||||
# create vector from token embeds
|
||||
token_embeds = self.text_encoder_list[text_encoder_idx].get_input_embeddings().weight.data
|
||||
# stack the tokens along batch axis adding that axis
|
||||
new_vector = torch.stack(
|
||||
[token_embeds[token_id] for token_id in self.placeholder_token_ids[text_encoder_idx]],
|
||||
dim=0
|
||||
)
|
||||
return new_vector
|
||||
|
||||
def _set_vec(self, new_vector, text_encoder_idx=0):
|
||||
# shape is (1, 768) for SD 1.5 for 1 token
|
||||
token_embeds = self.text_encoder_list[text_encoder_idx].get_input_embeddings().weight.data
|
||||
for i in range(new_vector.shape[0]):
|
||||
# apply the weights to the placeholder tokens while preserving gradient
|
||||
token_embeds[self.placeholder_token_ids[text_encoder_idx][i]] = new_vector[i].clone()
|
||||
|
||||
# make setter and getter for vec
|
||||
@property
|
||||
def vec(self):
|
||||
return self._get_vec(0)
|
||||
|
||||
@vec.setter
|
||||
def vec(self, new_vector):
|
||||
self._set_vec(new_vector, 0)
|
||||
|
||||
@property
|
||||
def vec2(self):
|
||||
return self._get_vec(1)
|
||||
|
||||
@vec2.setter
|
||||
def vec2(self, new_vector):
|
||||
self._set_vec(new_vector, 1)
|
||||
|
||||
# diffusers automatically expands the token meaning test123 becomes test123 test123_1 test123_2 etc
|
||||
# however, on training we don't use that pipeline, so we have to do it ourselves
|
||||
def inject_embedding_to_prompt(self, prompt, expand_token=False, to_replace_list=None, add_if_not_present=True):
|
||||
output_prompt = prompt
|
||||
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
|
||||
default_replacements = ["[name]", "[trigger]"]
|
||||
|
||||
replace_with = embedding_tokens if expand_token else self.trigger
|
||||
if to_replace_list is None:
|
||||
to_replace_list = default_replacements
|
||||
else:
|
||||
to_replace_list += default_replacements
|
||||
|
||||
# remove duplicates
|
||||
to_replace_list = list(set(to_replace_list))
|
||||
|
||||
# replace them all
|
||||
for to_replace in to_replace_list:
|
||||
# replace it
|
||||
output_prompt = output_prompt.replace(to_replace, replace_with)
|
||||
|
||||
# see how many times replace_with is in the prompt
|
||||
num_instances = output_prompt.count(replace_with)
|
||||
|
||||
if num_instances == 0 and add_if_not_present:
|
||||
# add it to the beginning of the prompt
|
||||
output_prompt = replace_with + " " + output_prompt
|
||||
|
||||
if num_instances > 1:
|
||||
print(
|
||||
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
|
||||
|
||||
return output_prompt
|
||||
|
||||
def state_dict(self):
|
||||
if self.sd.is_xl:
|
||||
state_dict = OrderedDict()
|
||||
state_dict['clip_l'] = self.vec
|
||||
state_dict['clip_g'] = self.vec2
|
||||
else:
|
||||
state_dict = OrderedDict()
|
||||
state_dict['emb_params'] = self.vec
|
||||
|
||||
return state_dict
|
||||
|
||||
def save(self, filename):
|
||||
# todo check to see how to get the vector out of the embedding
|
||||
|
||||
embedding_data = {
|
||||
"string_to_token": {"*": 265},
|
||||
"string_to_param": {"*": self.vec},
|
||||
"name": self.name,
|
||||
"step": self.step,
|
||||
# todo get these
|
||||
"sd_checkpoint": None,
|
||||
"sd_checkpoint_name": None,
|
||||
"notes": None,
|
||||
}
|
||||
# TODO we do not currently support this. Check how auto is doing it. Only safetensors supported sor sdxl
|
||||
if filename.endswith('.pt'):
|
||||
torch.save(embedding_data, filename)
|
||||
elif filename.endswith('.bin'):
|
||||
torch.save(embedding_data, filename)
|
||||
elif filename.endswith('.safetensors'):
|
||||
# save the embedding as a safetensors file
|
||||
state_dict = self.state_dict()
|
||||
# add all embedding data (except string_to_param), to metadata
|
||||
metadata = OrderedDict({k: json.dumps(v) for k, v in embedding_data.items() if k != "string_to_param"})
|
||||
metadata["string_to_param"] = {"*": "emb_params"}
|
||||
save_meta = get_meta_for_safetensors(metadata, name=self.name)
|
||||
save_file(state_dict, filename, metadata=save_meta)
|
||||
|
||||
def load_embedding_from_file(self, file_path, device):
|
||||
# full path
|
||||
path = os.path.realpath(file_path)
|
||||
filename = os.path.basename(path)
|
||||
name, ext = os.path.splitext(filename)
|
||||
tensors = {}
|
||||
ext = ext.upper()
|
||||
if ext in ['.PNG', '.WEBP', '.JXL', '.AVIF']:
|
||||
_, second_ext = os.path.splitext(name)
|
||||
if second_ext.upper() == '.PREVIEW':
|
||||
return
|
||||
|
||||
if ext in ['.BIN', '.PT']:
|
||||
# todo check this
|
||||
if self.sd.is_xl:
|
||||
raise Exception("XL not supported yet for bin, pt")
|
||||
data = torch.load(path, map_location="cpu")
|
||||
elif ext in ['.SAFETENSORS']:
|
||||
# rebuild the embedding from the safetensors file if it has it
|
||||
with safetensors.torch.safe_open(path, framework="pt", device="cpu") as f:
|
||||
metadata = f.metadata()
|
||||
for k in f.keys():
|
||||
tensors[k] = f.get_tensor(k)
|
||||
# data = safetensors.torch.load_file(path, device="cpu")
|
||||
if metadata and 'string_to_param' in metadata and 'emb_params' in tensors:
|
||||
# our format
|
||||
def try_json(v):
|
||||
try:
|
||||
return json.loads(v)
|
||||
except:
|
||||
return v
|
||||
|
||||
data = {k: try_json(v) for k, v in metadata.items()}
|
||||
data['string_to_param'] = {'*': tensors['emb_params']}
|
||||
else:
|
||||
# old format
|
||||
data = tensors
|
||||
else:
|
||||
return
|
||||
|
||||
if self.sd.is_xl:
|
||||
self.vec = tensors['clip_l'].detach().to(device, dtype=torch.float32)
|
||||
self.vec2 = tensors['clip_g'].detach().to(device, dtype=torch.float32)
|
||||
if 'step' in data:
|
||||
self.step = int(data['step'])
|
||||
else:
|
||||
# textual inversion embeddings
|
||||
if 'string_to_param' in data:
|
||||
param_dict = data['string_to_param']
|
||||
if hasattr(param_dict, '_parameters'):
|
||||
param_dict = getattr(param_dict,
|
||||
'_parameters') # fix for torch 1.12.1 loading saved file from torch 1.11
|
||||
assert len(param_dict) == 1, 'embedding file has multiple terms in it'
|
||||
emb = next(iter(param_dict.items()))[1]
|
||||
# diffuser concepts
|
||||
elif type(data) == dict and type(next(iter(data.values()))) == torch.Tensor:
|
||||
assert len(data.keys()) == 1, 'embedding file has multiple terms in it'
|
||||
|
||||
emb = next(iter(data.values()))
|
||||
if len(emb.shape) == 1:
|
||||
emb = emb.unsqueeze(0)
|
||||
else:
|
||||
raise Exception(
|
||||
f"Couldn't identify {filename} as neither textual inversion embedding nor diffuser concept.")
|
||||
|
||||
if 'step' in data:
|
||||
self.step = int(data['step'])
|
||||
|
||||
self.vec = emb.detach().to(device, dtype=torch.float32)
|
||||
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
|
||||
483
toolkit/image_utils.py
Normal file
483
toolkit/image_utils.py
Normal file
@@ -0,0 +1,483 @@
|
||||
# ref https://github.com/scardine/image_size/blob/master/get_image_size.py
|
||||
import atexit
|
||||
import collections
|
||||
import json
|
||||
import os
|
||||
import io
|
||||
import struct
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderTiny
|
||||
|
||||
FILE_UNKNOWN = "Sorry, don't know how to get size for this file."
|
||||
|
||||
|
||||
class UnknownImageFormat(Exception):
|
||||
pass
|
||||
|
||||
|
||||
types = collections.OrderedDict()
|
||||
BMP = types['BMP'] = 'BMP'
|
||||
GIF = types['GIF'] = 'GIF'
|
||||
ICO = types['ICO'] = 'ICO'
|
||||
JPEG = types['JPEG'] = 'JPEG'
|
||||
PNG = types['PNG'] = 'PNG'
|
||||
TIFF = types['TIFF'] = 'TIFF'
|
||||
|
||||
image_fields = ['path', 'type', 'file_size', 'width', 'height']
|
||||
|
||||
|
||||
class Image(collections.namedtuple('Image', image_fields)):
|
||||
|
||||
def to_str_row(self):
|
||||
return ("%d\t%d\t%d\t%s\t%s" % (
|
||||
self.width,
|
||||
self.height,
|
||||
self.file_size,
|
||||
self.type,
|
||||
self.path.replace('\t', '\\t'),
|
||||
))
|
||||
|
||||
def to_str_row_verbose(self):
|
||||
return ("%d\t%d\t%d\t%s\t%s\t##%s" % (
|
||||
self.width,
|
||||
self.height,
|
||||
self.file_size,
|
||||
self.type,
|
||||
self.path.replace('\t', '\\t'),
|
||||
self))
|
||||
|
||||
def to_str_json(self, indent=None):
|
||||
return json.dumps(self._asdict(), indent=indent)
|
||||
|
||||
|
||||
def get_image_size(file_path):
|
||||
"""
|
||||
Return (width, height) for a given img file content - no external
|
||||
dependencies except the os and struct builtin modules
|
||||
"""
|
||||
img = get_image_metadata(file_path)
|
||||
return (img.width, img.height)
|
||||
|
||||
|
||||
def get_image_size_from_bytesio(input, size):
|
||||
"""
|
||||
Return (width, height) for a given img file content - no external
|
||||
dependencies except the os and struct builtin modules
|
||||
|
||||
Args:
|
||||
input (io.IOBase): io object support read & seek
|
||||
size (int): size of buffer in byte
|
||||
"""
|
||||
img = get_image_metadata_from_bytesio(input, size)
|
||||
return (img.width, img.height)
|
||||
|
||||
|
||||
def get_image_metadata(file_path):
|
||||
"""
|
||||
Return an `Image` object for a given img file content - no external
|
||||
dependencies except the os and struct builtin modules
|
||||
|
||||
Args:
|
||||
file_path (str): path to an image file
|
||||
|
||||
Returns:
|
||||
Image: (path, type, file_size, width, height)
|
||||
"""
|
||||
size = os.path.getsize(file_path)
|
||||
|
||||
# be explicit with open arguments - we need binary mode
|
||||
with io.open(file_path, "rb") as input:
|
||||
return get_image_metadata_from_bytesio(input, size, file_path)
|
||||
|
||||
|
||||
def get_image_metadata_from_bytesio(input, size, file_path=None):
|
||||
"""
|
||||
Return an `Image` object for a given img file content - no external
|
||||
dependencies except the os and struct builtin modules
|
||||
|
||||
Args:
|
||||
input (io.IOBase): io object support read & seek
|
||||
size (int): size of buffer in byte
|
||||
file_path (str): path to an image file
|
||||
|
||||
Returns:
|
||||
Image: (path, type, file_size, width, height)
|
||||
"""
|
||||
height = -1
|
||||
width = -1
|
||||
data = input.read(26)
|
||||
msg = " raised while trying to decode as JPEG."
|
||||
|
||||
if (size >= 10) and data[:6] in (b'GIF87a', b'GIF89a'):
|
||||
# GIFs
|
||||
imgtype = GIF
|
||||
w, h = struct.unpack("<HH", data[6:10])
|
||||
width = int(w)
|
||||
height = int(h)
|
||||
elif ((size >= 24) and data.startswith(b'\211PNG\r\n\032\n')
|
||||
and (data[12:16] == b'IHDR')):
|
||||
# PNGs
|
||||
imgtype = PNG
|
||||
w, h = struct.unpack(">LL", data[16:24])
|
||||
width = int(w)
|
||||
height = int(h)
|
||||
elif (size >= 16) and data.startswith(b'\211PNG\r\n\032\n'):
|
||||
# older PNGs
|
||||
imgtype = PNG
|
||||
w, h = struct.unpack(">LL", data[8:16])
|
||||
width = int(w)
|
||||
height = int(h)
|
||||
elif (size >= 2) and data.startswith(b'\377\330'):
|
||||
# JPEG
|
||||
imgtype = JPEG
|
||||
input.seek(0)
|
||||
input.read(2)
|
||||
b = input.read(1)
|
||||
try:
|
||||
while (b and ord(b) != 0xDA):
|
||||
while (ord(b) != 0xFF):
|
||||
b = input.read(1)
|
||||
while (ord(b) == 0xFF):
|
||||
b = input.read(1)
|
||||
if (ord(b) >= 0xC0 and ord(b) <= 0xC3):
|
||||
input.read(3)
|
||||
h, w = struct.unpack(">HH", input.read(4))
|
||||
break
|
||||
else:
|
||||
input.read(
|
||||
int(struct.unpack(">H", input.read(2))[0]) - 2)
|
||||
b = input.read(1)
|
||||
width = int(w)
|
||||
height = int(h)
|
||||
except struct.error:
|
||||
raise UnknownImageFormat("StructError" + msg)
|
||||
except ValueError:
|
||||
raise UnknownImageFormat("ValueError" + msg)
|
||||
except Exception as e:
|
||||
raise UnknownImageFormat(e.__class__.__name__ + msg)
|
||||
elif (size >= 26) and data.startswith(b'BM'):
|
||||
# BMP
|
||||
imgtype = 'BMP'
|
||||
headersize = struct.unpack("<I", data[14:18])[0]
|
||||
if headersize == 12:
|
||||
w, h = struct.unpack("<HH", data[18:22])
|
||||
width = int(w)
|
||||
height = int(h)
|
||||
elif headersize >= 40:
|
||||
w, h = struct.unpack("<ii", data[18:26])
|
||||
width = int(w)
|
||||
# as h is negative when stored upside down
|
||||
height = abs(int(h))
|
||||
else:
|
||||
raise UnknownImageFormat(
|
||||
"Unkown DIB header size:" +
|
||||
str(headersize))
|
||||
elif (size >= 8) and data[:4] in (b"II\052\000", b"MM\000\052"):
|
||||
# Standard TIFF, big- or little-endian
|
||||
# BigTIFF and other different but TIFF-like formats are not
|
||||
# supported currently
|
||||
imgtype = TIFF
|
||||
byteOrder = data[:2]
|
||||
boChar = ">" if byteOrder == "MM" else "<"
|
||||
# maps TIFF type id to size (in bytes)
|
||||
# and python format char for struct
|
||||
tiffTypes = {
|
||||
1: (1, boChar + "B"), # BYTE
|
||||
2: (1, boChar + "c"), # ASCII
|
||||
3: (2, boChar + "H"), # SHORT
|
||||
4: (4, boChar + "L"), # LONG
|
||||
5: (8, boChar + "LL"), # RATIONAL
|
||||
6: (1, boChar + "b"), # SBYTE
|
||||
7: (1, boChar + "c"), # UNDEFINED
|
||||
8: (2, boChar + "h"), # SSHORT
|
||||
9: (4, boChar + "l"), # SLONG
|
||||
10: (8, boChar + "ll"), # SRATIONAL
|
||||
11: (4, boChar + "f"), # FLOAT
|
||||
12: (8, boChar + "d") # DOUBLE
|
||||
}
|
||||
ifdOffset = struct.unpack(boChar + "L", data[4:8])[0]
|
||||
try:
|
||||
countSize = 2
|
||||
input.seek(ifdOffset)
|
||||
ec = input.read(countSize)
|
||||
ifdEntryCount = struct.unpack(boChar + "H", ec)[0]
|
||||
# 2 bytes: TagId + 2 bytes: type + 4 bytes: count of values + 4
|
||||
# bytes: value offset
|
||||
ifdEntrySize = 12
|
||||
for i in range(ifdEntryCount):
|
||||
entryOffset = ifdOffset + countSize + i * ifdEntrySize
|
||||
input.seek(entryOffset)
|
||||
tag = input.read(2)
|
||||
tag = struct.unpack(boChar + "H", tag)[0]
|
||||
if (tag == 256 or tag == 257):
|
||||
# if type indicates that value fits into 4 bytes, value
|
||||
# offset is not an offset but value itself
|
||||
type = input.read(2)
|
||||
type = struct.unpack(boChar + "H", type)[0]
|
||||
if type not in tiffTypes:
|
||||
raise UnknownImageFormat(
|
||||
"Unkown TIFF field type:" +
|
||||
str(type))
|
||||
typeSize = tiffTypes[type][0]
|
||||
typeChar = tiffTypes[type][1]
|
||||
input.seek(entryOffset + 8)
|
||||
value = input.read(typeSize)
|
||||
value = int(struct.unpack(typeChar, value)[0])
|
||||
if tag == 256:
|
||||
width = value
|
||||
else:
|
||||
height = value
|
||||
if width > -1 and height > -1:
|
||||
break
|
||||
except Exception as e:
|
||||
raise UnknownImageFormat(str(e))
|
||||
elif size >= 2:
|
||||
# see http://en.wikipedia.org/wiki/ICO_(file_format)
|
||||
imgtype = 'ICO'
|
||||
input.seek(0)
|
||||
reserved = input.read(2)
|
||||
if 0 != struct.unpack("<H", reserved)[0]:
|
||||
raise UnknownImageFormat(FILE_UNKNOWN)
|
||||
format = input.read(2)
|
||||
assert 1 == struct.unpack("<H", format)[0]
|
||||
num = input.read(2)
|
||||
num = struct.unpack("<H", num)[0]
|
||||
if num > 1:
|
||||
import warnings
|
||||
warnings.warn("ICO File contains more than one image")
|
||||
# http://msdn.microsoft.com/en-us/library/ms997538.aspx
|
||||
w = input.read(1)
|
||||
h = input.read(1)
|
||||
width = ord(w)
|
||||
height = ord(h)
|
||||
else:
|
||||
raise UnknownImageFormat(FILE_UNKNOWN)
|
||||
|
||||
return Image(path=file_path,
|
||||
type=imgtype,
|
||||
file_size=size,
|
||||
width=width,
|
||||
height=height)
|
||||
|
||||
|
||||
import unittest
|
||||
|
||||
|
||||
class Test_get_image_size(unittest.TestCase):
|
||||
data = [{
|
||||
'path': 'lookmanodeps.png',
|
||||
'width': 251,
|
||||
'height': 208,
|
||||
'file_size': 22228,
|
||||
'type': 'PNG'}]
|
||||
|
||||
def setUp(self):
|
||||
pass
|
||||
|
||||
def test_get_image_size_from_bytesio(self):
|
||||
img = self.data[0]
|
||||
p = img['path']
|
||||
with io.open(p, 'rb') as fp:
|
||||
b = fp.read()
|
||||
fp = io.BytesIO(b)
|
||||
sz = len(b)
|
||||
output = get_image_size_from_bytesio(fp, sz)
|
||||
self.assertTrue(output)
|
||||
self.assertEqual(output,
|
||||
(img['width'],
|
||||
img['height']))
|
||||
|
||||
def test_get_image_metadata_from_bytesio(self):
|
||||
img = self.data[0]
|
||||
p = img['path']
|
||||
with io.open(p, 'rb') as fp:
|
||||
b = fp.read()
|
||||
fp = io.BytesIO(b)
|
||||
sz = len(b)
|
||||
output = get_image_metadata_from_bytesio(fp, sz)
|
||||
self.assertTrue(output)
|
||||
for field in image_fields:
|
||||
self.assertEqual(getattr(output, field), None if field == 'path' else img[field])
|
||||
|
||||
def test_get_image_metadata(self):
|
||||
img = self.data[0]
|
||||
output = get_image_metadata(img['path'])
|
||||
self.assertTrue(output)
|
||||
for field in image_fields:
|
||||
self.assertEqual(getattr(output, field), img[field])
|
||||
|
||||
def test_get_image_metadata__ENOENT_OSError(self):
|
||||
with self.assertRaises(OSError):
|
||||
get_image_metadata('THIS_DOES_NOT_EXIST')
|
||||
|
||||
def test_get_image_metadata__not_an_image_UnknownImageFormat(self):
|
||||
with self.assertRaises(UnknownImageFormat):
|
||||
get_image_metadata('README.rst')
|
||||
|
||||
def test_get_image_size(self):
|
||||
img = self.data[0]
|
||||
output = get_image_size(img['path'])
|
||||
self.assertTrue(output)
|
||||
self.assertEqual(output,
|
||||
(img['width'],
|
||||
img['height']))
|
||||
|
||||
def tearDown(self):
|
||||
pass
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
"""
|
||||
Print image metadata fields for the given file path.
|
||||
|
||||
Keyword Arguments:
|
||||
argv (list): commandline arguments (e.g. sys.argv[1:])
|
||||
Returns:
|
||||
int: zero for OK
|
||||
"""
|
||||
import logging
|
||||
import optparse
|
||||
import sys
|
||||
|
||||
prs = optparse.OptionParser(
|
||||
usage="%prog [-v|--verbose] [--json|--json-indent] <path0> [<pathN>]",
|
||||
description="Print metadata for the given image paths "
|
||||
"(without image library bindings).")
|
||||
|
||||
prs.add_option('--json',
|
||||
dest='json',
|
||||
action='store_true')
|
||||
prs.add_option('--json-indent',
|
||||
dest='json_indent',
|
||||
action='store_true')
|
||||
|
||||
prs.add_option('-v', '--verbose',
|
||||
dest='verbose',
|
||||
action='store_true', )
|
||||
prs.add_option('-q', '--quiet',
|
||||
dest='quiet',
|
||||
action='store_true', )
|
||||
prs.add_option('-t', '--test',
|
||||
dest='run_tests',
|
||||
action='store_true', )
|
||||
|
||||
argv = list(argv) if argv is not None else sys.argv[1:]
|
||||
(opts, args) = prs.parse_args(args=argv)
|
||||
loglevel = logging.INFO
|
||||
if opts.verbose:
|
||||
loglevel = logging.DEBUG
|
||||
elif opts.quiet:
|
||||
loglevel = logging.ERROR
|
||||
logging.basicConfig(level=loglevel)
|
||||
log = logging.getLogger()
|
||||
log.debug('argv: %r', argv)
|
||||
log.debug('opts: %r', opts)
|
||||
log.debug('args: %r', args)
|
||||
|
||||
if opts.run_tests:
|
||||
import sys
|
||||
sys.argv = [sys.argv[0]] + args
|
||||
import unittest
|
||||
return unittest.main()
|
||||
|
||||
output_func = Image.to_str_row
|
||||
if opts.json_indent:
|
||||
import functools
|
||||
output_func = functools.partial(Image.to_str_json, indent=2)
|
||||
elif opts.json:
|
||||
output_func = Image.to_str_json
|
||||
elif opts.verbose:
|
||||
output_func = Image.to_str_row_verbose
|
||||
|
||||
EX_OK = 0
|
||||
EX_NOT_OK = 2
|
||||
|
||||
if len(args) < 1:
|
||||
prs.print_help()
|
||||
print('')
|
||||
prs.error("You must specify one or more paths to image files")
|
||||
|
||||
errors = []
|
||||
for path_arg in args:
|
||||
try:
|
||||
img = get_image_metadata(path_arg)
|
||||
print(output_func(img))
|
||||
except KeyboardInterrupt:
|
||||
raise
|
||||
except OSError as e:
|
||||
log.error((path_arg, e))
|
||||
errors.append((path_arg, e))
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
errors.append((path_arg, e))
|
||||
pass
|
||||
if len(errors):
|
||||
import pprint
|
||||
print("ERRORS", file=sys.stderr)
|
||||
print("======", file=sys.stderr)
|
||||
print(pprint.pformat(errors, indent=2), file=sys.stderr)
|
||||
return EX_NOT_OK
|
||||
return EX_OK
|
||||
|
||||
|
||||
is_window_shown = False
|
||||
|
||||
|
||||
def show_img(img, name='AI Toolkit'):
|
||||
global is_window_shown
|
||||
|
||||
img = np.clip(img, 0, 255).astype(np.uint8)
|
||||
cv2.imshow(name, img[:, :, ::-1])
|
||||
k = cv2.waitKey(10) & 0xFF
|
||||
if k == 27: # Esc key to stop
|
||||
print('\nESC pressed, stopping')
|
||||
raise KeyboardInterrupt
|
||||
if not is_window_shown:
|
||||
is_window_shown = True
|
||||
|
||||
|
||||
|
||||
def show_tensors(imgs: torch.Tensor, name='AI Toolkit'):
|
||||
# if rank is 4
|
||||
if len(imgs.shape) == 4:
|
||||
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
|
||||
else:
|
||||
img_list = [imgs]
|
||||
# put images side by side
|
||||
img = torch.cat(img_list, dim=3)
|
||||
# img is -1 to 1, convert to 0 to 255
|
||||
img = img / 2 + 0.5
|
||||
img_numpy = img.to(torch.float32).detach().cpu().numpy()
|
||||
img_numpy = np.clip(img_numpy, 0, 1) * 255
|
||||
# convert to numpy Move channel to last
|
||||
img_numpy = img_numpy.transpose(0, 2, 3, 1)
|
||||
# convert to uint8
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
show_img(img_numpy[0], name=name)
|
||||
|
||||
|
||||
def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit'):
|
||||
# decode latents
|
||||
if vae.device == 'cpu':
|
||||
vae.to(latents.device)
|
||||
latents = latents / vae.config['scaling_factor']
|
||||
imgs = vae.decode(latents).sample
|
||||
show_tensors(imgs, name=name)
|
||||
|
||||
|
||||
|
||||
def on_exit():
|
||||
if is_window_shown:
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
|
||||
atexit.register(on_exit)
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(main(argv=sys.argv[1:]))
|
||||
410
toolkit/inversion_utils.py
Normal file
410
toolkit/inversion_utils.py
Normal file
@@ -0,0 +1,410 @@
|
||||
# ref https://huggingface.co/spaces/editing-images/ledits/blob/main/inversion_utils.py
|
||||
|
||||
import torch
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit import train_tools
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
def mu_tilde(model, xt, x0, timestep):
|
||||
"mu_tilde(x_t, x_0) DDPM paper eq. 7"
|
||||
prev_timestep = timestep - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps
|
||||
alpha_prod_t_prev = model.scheduler.alphas_cumprod[
|
||||
prev_timestep] if prev_timestep >= 0 else model.scheduler.final_alpha_cumprod
|
||||
alpha_t = model.scheduler.alphas[timestep]
|
||||
beta_t = 1 - alpha_t
|
||||
alpha_bar = model.scheduler.alphas_cumprod[timestep]
|
||||
return ((alpha_prod_t_prev ** 0.5 * beta_t) / (1 - alpha_bar)) * x0 + (
|
||||
(alpha_t ** 0.5 * (1 - alpha_prod_t_prev)) / (1 - alpha_bar)) * xt
|
||||
|
||||
|
||||
def sample_xts_from_x0(sd: StableDiffusion, sample: torch.Tensor, num_inference_steps=50):
|
||||
"""
|
||||
Samples from P(x_1:T|x_0)
|
||||
"""
|
||||
# torch.manual_seed(43256465436)
|
||||
alpha_bar = sd.noise_scheduler.alphas_cumprod
|
||||
sqrt_one_minus_alpha_bar = (1 - alpha_bar) ** 0.5
|
||||
alphas = sd.noise_scheduler.alphas
|
||||
betas = 1 - alphas
|
||||
# variance_noise_shape = (
|
||||
# num_inference_steps,
|
||||
# sd.unet.in_channels,
|
||||
# sd.unet.sample_size,
|
||||
# sd.unet.sample_size)
|
||||
variance_noise_shape = list(sample.shape)
|
||||
variance_noise_shape[0] = num_inference_steps
|
||||
|
||||
timesteps = sd.noise_scheduler.timesteps.to(sd.device)
|
||||
t_to_idx = {int(v): k for k, v in enumerate(timesteps)}
|
||||
xts = torch.zeros(variance_noise_shape).to(sample.device, dtype=torch.float16)
|
||||
for t in reversed(timesteps):
|
||||
idx = t_to_idx[int(t)]
|
||||
xts[idx] = sample * (alpha_bar[t] ** 0.5) + torch.randn_like(sample, dtype=torch.float16) * sqrt_one_minus_alpha_bar[t]
|
||||
xts = torch.cat([xts, sample], dim=0)
|
||||
|
||||
return xts
|
||||
|
||||
|
||||
def encode_text(model, prompts):
|
||||
text_input = model.tokenizer(
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=model.tokenizer.model_max_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
with torch.no_grad():
|
||||
text_encoding = model.text_encoder(text_input.input_ids.to(model.device))[0]
|
||||
return text_encoding
|
||||
|
||||
|
||||
def forward_step(sd: StableDiffusion, model_output, timestep, sample):
|
||||
next_timestep = min(
|
||||
sd.noise_scheduler.config['num_train_timesteps'] - 2,
|
||||
timestep + sd.noise_scheduler.config['num_train_timesteps'] // sd.noise_scheduler.num_inference_steps
|
||||
)
|
||||
|
||||
# 2. compute alphas, betas
|
||||
alpha_prod_t = sd.noise_scheduler.alphas_cumprod[timestep]
|
||||
# alpha_prod_t_next = self.scheduler.alphas_cumprod[next_timestep] if next_ltimestep >= 0 else self.scheduler.final_alpha_cumprod
|
||||
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
# 3. compute predicted original sample from predicted noise also called
|
||||
# "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf
|
||||
pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)
|
||||
|
||||
# 5. TODO: simple noising implementation
|
||||
next_sample = sd.noise_scheduler.add_noise(
|
||||
pred_original_sample,
|
||||
model_output,
|
||||
torch.LongTensor([next_timestep]))
|
||||
return next_sample
|
||||
|
||||
|
||||
def get_variance(sd: StableDiffusion, timestep): # , prev_timestep):
|
||||
prev_timestep = timestep - sd.noise_scheduler.config['num_train_timesteps'] // sd.noise_scheduler.num_inference_steps
|
||||
alpha_prod_t = sd.noise_scheduler.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = sd.noise_scheduler.alphas_cumprod[
|
||||
prev_timestep] if prev_timestep >= 0 else sd.noise_scheduler.final_alpha_cumprod
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
beta_prod_t_prev = 1 - alpha_prod_t_prev
|
||||
variance = (beta_prod_t_prev / beta_prod_t) * (1 - alpha_prod_t / alpha_prod_t_prev)
|
||||
return variance
|
||||
|
||||
|
||||
def get_time_ids_from_latents(sd: StableDiffusion, latents: torch.Tensor):
|
||||
VAE_SCALE_FACTOR = 2 ** (len(sd.vae.config['block_out_channels']) - 1)
|
||||
if sd.is_xl:
|
||||
bs, ch, h, w = list(latents.shape)
|
||||
|
||||
height = h * VAE_SCALE_FACTOR
|
||||
width = w * VAE_SCALE_FACTOR
|
||||
|
||||
dtype = latents.dtype
|
||||
# just do it without any cropping nonsense
|
||||
target_size = (height, width)
|
||||
original_size = (height, width)
|
||||
crops_coords_top_left = (0, 0)
|
||||
add_time_ids = list(original_size + crops_coords_top_left + target_size)
|
||||
add_time_ids = torch.tensor([add_time_ids])
|
||||
add_time_ids = add_time_ids.to(latents.device, dtype=dtype)
|
||||
|
||||
batch_time_ids = torch.cat(
|
||||
[add_time_ids for _ in range(bs)]
|
||||
)
|
||||
return batch_time_ids
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
def inversion_forward_process(
|
||||
sd: StableDiffusion,
|
||||
sample: torch.Tensor,
|
||||
conditional_embeddings: PromptEmbeds,
|
||||
unconditional_embeddings: PromptEmbeds,
|
||||
etas=None,
|
||||
prog_bar=False,
|
||||
cfg_scale=3.5,
|
||||
num_inference_steps=50, eps=None
|
||||
):
|
||||
current_num_timesteps = len(sd.noise_scheduler.timesteps)
|
||||
sd.noise_scheduler.set_timesteps(num_inference_steps, device=sd.device)
|
||||
|
||||
timesteps = sd.noise_scheduler.timesteps.to(sd.device)
|
||||
# variance_noise_shape = (
|
||||
# num_inference_steps,
|
||||
# sd.unet.in_channels,
|
||||
# sd.unet.sample_size,
|
||||
# sd.unet.sample_size
|
||||
# )
|
||||
variance_noise_shape = list(sample.shape)
|
||||
variance_noise_shape[0] = num_inference_steps
|
||||
if etas is None or (type(etas) in [int, float] and etas == 0):
|
||||
eta_is_zero = True
|
||||
zs = None
|
||||
else:
|
||||
eta_is_zero = False
|
||||
if type(etas) in [int, float]: etas = [etas] * sd.noise_scheduler.num_inference_steps
|
||||
xts = sample_xts_from_x0(sd, sample, num_inference_steps=num_inference_steps)
|
||||
alpha_bar = sd.noise_scheduler.alphas_cumprod
|
||||
zs = torch.zeros(size=variance_noise_shape, device=sd.device, dtype=torch.float16)
|
||||
|
||||
t_to_idx = {int(v): k for k, v in enumerate(timesteps)}
|
||||
noisy_sample = sample
|
||||
op = tqdm(reversed(timesteps), desc="Inverting...") if prog_bar else reversed(timesteps)
|
||||
|
||||
for timestep in op:
|
||||
idx = t_to_idx[int(timestep)]
|
||||
# 1. predict noise residual
|
||||
if not eta_is_zero:
|
||||
noisy_sample = xts[idx][None]
|
||||
|
||||
added_cond_kwargs = {}
|
||||
|
||||
with torch.no_grad():
|
||||
text_embeddings = train_tools.concat_prompt_embeddings(
|
||||
unconditional_embeddings, # negative embedding
|
||||
conditional_embeddings, # positive embedding
|
||||
1, # batch size
|
||||
)
|
||||
if sd.is_xl:
|
||||
add_time_ids = get_time_ids_from_latents(sd, noisy_sample)
|
||||
# add extra for cfg
|
||||
add_time_ids = torch.cat(
|
||||
[add_time_ids] * 2, dim=0
|
||||
)
|
||||
|
||||
added_cond_kwargs = {
|
||||
"text_embeds": text_embeddings.pooled_embeds,
|
||||
"time_ids": add_time_ids,
|
||||
}
|
||||
|
||||
# double up for cfg
|
||||
latent_model_input = torch.cat(
|
||||
[noisy_sample] * 2, dim=0
|
||||
)
|
||||
|
||||
noise_pred = sd.unet(
|
||||
latent_model_input,
|
||||
timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
).sample
|
||||
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
|
||||
# out = sd.unet.forward(noisy_sample, timestep=timestep, encoder_hidden_states=uncond_embedding)
|
||||
# cond_out = sd.unet.forward(noisy_sample, timestep=timestep, encoder_hidden_states=text_embeddings)
|
||||
|
||||
noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
if eta_is_zero:
|
||||
# 2. compute more noisy image and set x_t -> x_t+1
|
||||
noisy_sample = forward_step(sd, noise_pred, timestep, noisy_sample)
|
||||
xts = None
|
||||
|
||||
else:
|
||||
xtm1 = xts[idx + 1][None]
|
||||
# pred of x0
|
||||
pred_original_sample = (noisy_sample - (1 - alpha_bar[timestep]) ** 0.5 * noise_pred) / alpha_bar[
|
||||
timestep] ** 0.5
|
||||
|
||||
# direction to xt
|
||||
prev_timestep = timestep - sd.noise_scheduler.config[
|
||||
'num_train_timesteps'] // sd.noise_scheduler.num_inference_steps
|
||||
alpha_prod_t_prev = sd.noise_scheduler.alphas_cumprod[
|
||||
prev_timestep] if prev_timestep >= 0 else sd.noise_scheduler.final_alpha_cumprod
|
||||
|
||||
variance = get_variance(sd, timestep)
|
||||
pred_sample_direction = (1 - alpha_prod_t_prev - etas[idx] * variance) ** (0.5) * noise_pred
|
||||
|
||||
mu_xt = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction
|
||||
|
||||
z = (xtm1 - mu_xt) / (etas[idx] * variance ** 0.5)
|
||||
zs[idx] = z
|
||||
|
||||
# correction to avoid error accumulation
|
||||
xtm1 = mu_xt + (etas[idx] * variance ** 0.5) * z
|
||||
xts[idx + 1] = xtm1
|
||||
|
||||
if not zs is None:
|
||||
zs[-1] = torch.zeros_like(zs[-1])
|
||||
|
||||
# restore timesteps
|
||||
sd.noise_scheduler.set_timesteps(current_num_timesteps, device=sd.device)
|
||||
|
||||
return noisy_sample, zs, xts
|
||||
|
||||
|
||||
#
|
||||
# def inversion_forward_process(
|
||||
# model,
|
||||
# sample,
|
||||
# etas=None,
|
||||
# prog_bar=False,
|
||||
# prompt="",
|
||||
# cfg_scale=3.5,
|
||||
# num_inference_steps=50, eps=None
|
||||
# ):
|
||||
# if not prompt == "":
|
||||
# text_embeddings = encode_text(model, prompt)
|
||||
# uncond_embedding = encode_text(model, "")
|
||||
# timesteps = model.scheduler.timesteps.to(model.device)
|
||||
# variance_noise_shape = (
|
||||
# num_inference_steps,
|
||||
# model.unet.in_channels,
|
||||
# model.unet.sample_size,
|
||||
# model.unet.sample_size)
|
||||
# if etas is None or (type(etas) in [int, float] and etas == 0):
|
||||
# eta_is_zero = True
|
||||
# zs = None
|
||||
# else:
|
||||
# eta_is_zero = False
|
||||
# if type(etas) in [int, float]: etas = [etas] * model.scheduler.num_inference_steps
|
||||
# xts = sample_xts_from_x0(model, sample, num_inference_steps=num_inference_steps)
|
||||
# alpha_bar = model.scheduler.alphas_cumprod
|
||||
# zs = torch.zeros(size=variance_noise_shape, device=model.device, dtype=torch.float16)
|
||||
#
|
||||
# t_to_idx = {int(v): k for k, v in enumerate(timesteps)}
|
||||
# noisy_sample = sample
|
||||
# op = tqdm(reversed(timesteps), desc="Inverting...") if prog_bar else reversed(timesteps)
|
||||
#
|
||||
# for t in op:
|
||||
# idx = t_to_idx[int(t)]
|
||||
# # 1. predict noise residual
|
||||
# if not eta_is_zero:
|
||||
# noisy_sample = xts[idx][None]
|
||||
#
|
||||
# with torch.no_grad():
|
||||
# out = model.unet.forward(noisy_sample, timestep=t, encoder_hidden_states=uncond_embedding)
|
||||
# if not prompt == "":
|
||||
# cond_out = model.unet.forward(noisy_sample, timestep=t, encoder_hidden_states=text_embeddings)
|
||||
#
|
||||
# if not prompt == "":
|
||||
# ## classifier free guidance
|
||||
# noise_pred = out.sample + cfg_scale * (cond_out.sample - out.sample)
|
||||
# else:
|
||||
# noise_pred = out.sample
|
||||
#
|
||||
# if eta_is_zero:
|
||||
# # 2. compute more noisy image and set x_t -> x_t+1
|
||||
# noisy_sample = forward_step(model, noise_pred, t, noisy_sample)
|
||||
#
|
||||
# else:
|
||||
# xtm1 = xts[idx + 1][None]
|
||||
# # pred of x0
|
||||
# pred_original_sample = (noisy_sample - (1 - alpha_bar[t]) ** 0.5 * noise_pred) / alpha_bar[t] ** 0.5
|
||||
#
|
||||
# # direction to xt
|
||||
# prev_timestep = t - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps
|
||||
# alpha_prod_t_prev = model.scheduler.alphas_cumprod[
|
||||
# prev_timestep] if prev_timestep >= 0 else model.scheduler.final_alpha_cumprod
|
||||
#
|
||||
# variance = get_variance(model, t)
|
||||
# pred_sample_direction = (1 - alpha_prod_t_prev - etas[idx] * variance) ** (0.5) * noise_pred
|
||||
#
|
||||
# mu_xt = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction
|
||||
#
|
||||
# z = (xtm1 - mu_xt) / (etas[idx] * variance ** 0.5)
|
||||
# zs[idx] = z
|
||||
#
|
||||
# # correction to avoid error accumulation
|
||||
# xtm1 = mu_xt + (etas[idx] * variance ** 0.5) * z
|
||||
# xts[idx + 1] = xtm1
|
||||
#
|
||||
# if not zs is None:
|
||||
# zs[-1] = torch.zeros_like(zs[-1])
|
||||
#
|
||||
# return noisy_sample, zs, xts
|
||||
|
||||
|
||||
def reverse_step(model, model_output, timestep, sample, eta=0, variance_noise=None):
|
||||
# 1. get previous step value (=t-1)
|
||||
prev_timestep = timestep - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps
|
||||
# 2. compute alphas, betas
|
||||
alpha_prod_t = model.scheduler.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = model.scheduler.alphas_cumprod[
|
||||
prev_timestep] if prev_timestep >= 0 else model.scheduler.final_alpha_cumprod
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
# 3. compute predicted original sample from predicted noise also called
|
||||
# "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf
|
||||
pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)
|
||||
# 5. compute variance: "sigma_t(η)" -> see formula (16)
|
||||
# σ_t = sqrt((1 − α_t−1)/(1 − α_t)) * sqrt(1 − α_t/α_t−1)
|
||||
# variance = self.scheduler._get_variance(timestep, prev_timestep)
|
||||
variance = get_variance(model, timestep) # , prev_timestep)
|
||||
std_dev_t = eta * variance ** (0.5)
|
||||
# Take care of asymetric reverse process (asyrp)
|
||||
model_output_direction = model_output
|
||||
# 6. compute "direction pointing to x_t" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf
|
||||
# pred_sample_direction = (1 - alpha_prod_t_prev - std_dev_t**2) ** (0.5) * model_output_direction
|
||||
pred_sample_direction = (1 - alpha_prod_t_prev - eta * variance) ** (0.5) * model_output_direction
|
||||
# 7. compute x_t without "random noise" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf
|
||||
prev_sample = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction
|
||||
# 8. Add noice if eta > 0
|
||||
if eta > 0:
|
||||
if variance_noise is None:
|
||||
variance_noise = torch.randn(model_output.shape, device=model.device, dtype=torch.float16)
|
||||
sigma_z = eta * variance ** (0.5) * variance_noise
|
||||
prev_sample = prev_sample + sigma_z
|
||||
|
||||
return prev_sample
|
||||
|
||||
|
||||
def inversion_reverse_process(
|
||||
model,
|
||||
xT,
|
||||
etas=0,
|
||||
prompts="",
|
||||
cfg_scales=None,
|
||||
prog_bar=False,
|
||||
zs=None,
|
||||
controller=None,
|
||||
asyrp=False):
|
||||
batch_size = len(prompts)
|
||||
|
||||
cfg_scales_tensor = torch.Tensor(cfg_scales).view(-1, 1, 1, 1).to(model.device, dtype=torch.float16)
|
||||
|
||||
text_embeddings = encode_text(model, prompts)
|
||||
uncond_embedding = encode_text(model, [""] * batch_size)
|
||||
|
||||
if etas is None: etas = 0
|
||||
if type(etas) in [int, float]: etas = [etas] * model.scheduler.num_inference_steps
|
||||
assert len(etas) == model.scheduler.num_inference_steps
|
||||
timesteps = model.scheduler.timesteps.to(model.device)
|
||||
|
||||
xt = xT.expand(batch_size, -1, -1, -1)
|
||||
op = tqdm(timesteps[-zs.shape[0]:]) if prog_bar else timesteps[-zs.shape[0]:]
|
||||
|
||||
t_to_idx = {int(v): k for k, v in enumerate(timesteps[-zs.shape[0]:])}
|
||||
|
||||
for t in op:
|
||||
idx = t_to_idx[int(t)]
|
||||
## Unconditional embedding
|
||||
with torch.no_grad():
|
||||
uncond_out = model.unet.forward(xt, timestep=t,
|
||||
encoder_hidden_states=uncond_embedding)
|
||||
|
||||
## Conditional embedding
|
||||
if prompts:
|
||||
with torch.no_grad():
|
||||
cond_out = model.unet.forward(xt, timestep=t,
|
||||
encoder_hidden_states=text_embeddings)
|
||||
|
||||
z = zs[idx] if not zs is None else None
|
||||
z = z.expand(batch_size, -1, -1, -1)
|
||||
if prompts:
|
||||
## classifier free guidance
|
||||
noise_pred = uncond_out.sample + cfg_scales_tensor * (cond_out.sample - uncond_out.sample)
|
||||
else:
|
||||
noise_pred = uncond_out.sample
|
||||
# 2. compute less noisy image and set x_t -> x_t-1
|
||||
xt = reverse_step(model, noise_pred, t, xt, eta=etas[idx], variance_noise=z)
|
||||
if controller is not None:
|
||||
xt = controller.step_callback(xt)
|
||||
return xt, zs
|
||||
157
toolkit/ip_adapter.py
Normal file
157
toolkit/ip_adapter.py
Normal file
@@ -0,0 +1,157 @@
|
||||
import torch
|
||||
import sys
|
||||
|
||||
from PIL import Image
|
||||
from torch.nn import Parameter
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from typing import TYPE_CHECKING, Union, Iterator, Mapping, Any, Tuple, List
|
||||
from collections import OrderedDict
|
||||
from ipadapter.ip_adapter.attention_processor import AttnProcessor, IPAttnProcessor
|
||||
from ipadapter.ip_adapter.ip_adapter import ImageProjModel
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
import weakref
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
# loosely based on # ref https://github.com/tencent-ailab/IP-Adapter/blob/main/tutorial_train.py
|
||||
class IPAdapter(torch.nn.Module):
|
||||
"""IP-Adapter"""
|
||||
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: 'AdapterConfig'):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.clip_image_processor = CLIPImageProcessor()
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(adapter_config.image_encoder_path)
|
||||
if adapter_config.type == 'ip':
|
||||
# ip-adapter
|
||||
image_proj_model = ImageProjModel(
|
||||
cross_attention_dim=sd.unet.config['cross_attention_dim'],
|
||||
clip_embeddings_dim=self.image_encoder.config.projection_dim,
|
||||
clip_extra_context_tokens=4,
|
||||
)
|
||||
elif adapter_config.type == 'ip+':
|
||||
# ip-adapter-plus
|
||||
num_tokens = 16
|
||||
image_proj_model = Resampler(
|
||||
dim=sd.unet.config['cross_attention_dim'],
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=12,
|
||||
num_queries=num_tokens,
|
||||
embedding_dim=self.image_encoder.config.hidden_size,
|
||||
output_dim=sd.unet.config['cross_attention_dim'],
|
||||
ff_mult=4
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown adapter type: {adapter_config.type}")
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
for name in sd.unet.attn_processors.keys():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else sd.unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = sd.unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(sd.unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AttnProcessor()
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
weights = {
|
||||
"to_k_ip.weight": unet_sd[layer_name + ".to_k.weight"],
|
||||
"to_v_ip.weight": unet_sd[layer_name + ".to_v.weight"],
|
||||
}
|
||||
attn_procs[name] = IPAttnProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim)
|
||||
attn_procs[name].load_state_dict(weights)
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
|
||||
sd.adapter = self
|
||||
self.unet_ref: weakref.ref = weakref.ref(sd.unet)
|
||||
self.image_proj_model = image_proj_model
|
||||
self.adapter_modules = adapter_modules
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
super().to(*args, **kwargs)
|
||||
self.image_encoder.to(*args, **kwargs)
|
||||
self.image_proj_model.to(*args, **kwargs)
|
||||
self.adapter_modules.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def load_ip_adapter(self, state_dict: Union[OrderedDict, dict]):
|
||||
self.image_proj_model.load_state_dict(state_dict["image_proj"])
|
||||
ip_layers = torch.nn.ModuleList(self.pipe.unet.attn_processors.values())
|
||||
ip_layers.load_state_dict(state_dict["ip_adapter"])
|
||||
|
||||
def state_dict(self) -> OrderedDict:
|
||||
state_dict = OrderedDict()
|
||||
state_dict["image_proj"] = self.image_proj_model.state_dict()
|
||||
state_dict["ip_adapter"] = self.adapter_modules.state_dict()
|
||||
return state_dict
|
||||
|
||||
def set_scale(self, scale):
|
||||
for attn_processor in self.pipe.unet.attn_processors.values():
|
||||
if isinstance(attn_processor, IPAttnProcessor):
|
||||
attn_processor.scale = scale
|
||||
|
||||
@torch.no_grad()
|
||||
def get_clip_image_embeds_from_pil(self, pil_image: Union[Image.Image, List[Image.Image]], drop=False) -> torch.Tensor:
|
||||
# todo: add support for sdxl
|
||||
if isinstance(pil_image, Image.Image):
|
||||
pil_image = [pil_image]
|
||||
clip_image = self.clip_image_processor(images=pil_image, return_tensors="pt").pixel_values
|
||||
clip_image = clip_image.to(self.device, dtype=torch.float16)
|
||||
if drop:
|
||||
clip_image = clip_image * 0
|
||||
clip_image_embeds = self.image_encoder(clip_image, output_hidden_states=True).hidden_states[-2]
|
||||
return clip_image_embeds
|
||||
|
||||
@torch.no_grad()
|
||||
def get_clip_image_embeds_from_tensors(self, tensors_0_1: torch.Tensor, drop=False) -> torch.Tensor:
|
||||
# tensors should be 0-1
|
||||
# todo: add support for sdxl
|
||||
if tensors_0_1.ndim == 3:
|
||||
tensors_0_1 = tensors_0_1.unsqueeze(0)
|
||||
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
|
||||
clip_image = self.clip_image_processor(images=tensors_0_1, return_tensors="pt", do_resize=False).pixel_values
|
||||
clip_image = clip_image.to(self.device, dtype=torch.float16)
|
||||
if drop:
|
||||
clip_image = clip_image * 0
|
||||
clip_image_embeds = self.image_encoder(clip_image, output_hidden_states=True).hidden_states[-2]
|
||||
return clip_image_embeds
|
||||
|
||||
# use drop for prompt dropout, or negatives
|
||||
def forward(self, embeddings: PromptEmbeds, clip_image_embeds: torch.Tensor) -> PromptEmbeds:
|
||||
clip_image_embeds = clip_image_embeds.detach()
|
||||
clip_image_embeds = clip_image_embeds.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
image_prompt_embeds = self.image_proj_model(clip_image_embeds.detach())
|
||||
embeddings.text_embeds = torch.cat([embeddings.text_embeds, image_prompt_embeds], dim=1)
|
||||
return embeddings
|
||||
|
||||
def parameters(self, recurse: bool = True) -> Iterator[Parameter]:
|
||||
for attn_processor in self.adapter_modules:
|
||||
yield from attn_processor.parameters(recurse)
|
||||
yield from self.image_proj_model.parameters(recurse)
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||
self.image_proj_model.load_state_dict(state_dict["image_proj"], strict=strict)
|
||||
self.adapter_modules.load_state_dict(state_dict["ip_adapter"], strict=strict)
|
||||
@@ -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()
|
||||
|
||||
3154
toolkit/keymaps/stable_diffusion_locon_sdxl.json
Normal file
3154
toolkit/keymaps/stable_diffusion_locon_sdxl.json
Normal file
File diff suppressed because it is too large
Load Diff
3498
toolkit/keymaps/stable_diffusion_refiner.json
Normal file
3498
toolkit/keymaps/stable_diffusion_refiner.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_refiner_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_refiner_ldm_base.safetensors
Normal file
Binary file not shown.
27
toolkit/keymaps/stable_diffusion_refiner_unmatched.json
Normal file
27
toolkit/keymaps/stable_diffusion_refiner_unmatched.json
Normal file
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"ldm": {
|
||||
"conditioner.embedders.0.model.logit_scale": {
|
||||
"shape": [],
|
||||
"min": 4.60546875,
|
||||
"max": 4.60546875
|
||||
},
|
||||
"conditioner.embedders.0.model.text_projection": {
|
||||
"shape": [
|
||||
1280,
|
||||
1280
|
||||
],
|
||||
"min": -0.15966796875,
|
||||
"max": 0.230712890625
|
||||
}
|
||||
},
|
||||
"diffusers": {
|
||||
"te1_text_projection.weight": {
|
||||
"shape": [
|
||||
1280,
|
||||
1280
|
||||
],
|
||||
"min": -0.15966796875,
|
||||
"max": 0.230712890625
|
||||
}
|
||||
}
|
||||
}
|
||||
1234
toolkit/keymaps/stable_diffusion_sd1.json
Normal file
1234
toolkit/keymaps/stable_diffusion_sd1.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_sd1_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_sd1_ldm_base.safetensors
Normal file
Binary file not shown.
2424
toolkit/keymaps/stable_diffusion_sd2.json
Normal file
2424
toolkit/keymaps/stable_diffusion_sd2.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_sd2_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_sd2_ldm_base.safetensors
Normal file
Binary file not shown.
200
toolkit/keymaps/stable_diffusion_sd2_unmatched.json
Normal file
200
toolkit/keymaps/stable_diffusion_sd2_unmatched.json
Normal file
@@ -0,0 +1,200 @@
|
||||
{
|
||||
"ldm": {
|
||||
"alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.00466156005859375,
|
||||
"max": 0.9990234375
|
||||
},
|
||||
"alphas_cumprod_prev": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0047149658203125,
|
||||
"max": 1.0
|
||||
},
|
||||
"betas": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0008502006530761719,
|
||||
"max": 0.01200103759765625
|
||||
},
|
||||
"cond_stage_model.model.logit_scale": {
|
||||
"shape": [],
|
||||
"min": 4.60546875,
|
||||
"max": 4.60546875
|
||||
},
|
||||
"cond_stage_model.model.text_projection": {
|
||||
"shape": [
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
"min": -0.109130859375,
|
||||
"max": 0.09271240234375
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.attn.in_proj_bias": {
|
||||
"shape": [
|
||||
3072
|
||||
],
|
||||
"min": -2.525390625,
|
||||
"max": 2.591796875
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.attn.in_proj_weight": {
|
||||
"shape": [
|
||||
3072,
|
||||
1024
|
||||
],
|
||||
"min": -0.12261962890625,
|
||||
"max": 0.1258544921875
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.attn.out_proj.bias": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": -0.422607421875,
|
||||
"max": 1.17578125
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.attn.out_proj.weight": {
|
||||
"shape": [
|
||||
1024,
|
||||
1024
|
||||
],
|
||||
"min": -0.0738525390625,
|
||||
"max": 0.08673095703125
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.ln_1.bias": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": -3.392578125,
|
||||
"max": 0.90625
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.ln_1.weight": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": 0.379638671875,
|
||||
"max": 2.02734375
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.ln_2.bias": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": -0.833984375,
|
||||
"max": 2.525390625
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.ln_2.weight": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": 1.17578125,
|
||||
"max": 2.037109375
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.mlp.c_fc.bias": {
|
||||
"shape": [
|
||||
4096
|
||||
],
|
||||
"min": -1.619140625,
|
||||
"max": 0.5595703125
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.mlp.c_fc.weight": {
|
||||
"shape": [
|
||||
4096,
|
||||
1024
|
||||
],
|
||||
"min": -0.08953857421875,
|
||||
"max": 0.13232421875
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.mlp.c_proj.bias": {
|
||||
"shape": [
|
||||
1024
|
||||
],
|
||||
"min": -1.8662109375,
|
||||
"max": 0.74658203125
|
||||
},
|
||||
"cond_stage_model.model.transformer.resblocks.23.mlp.c_proj.weight": {
|
||||
"shape": [
|
||||
1024,
|
||||
4096
|
||||
],
|
||||
"min": -0.12939453125,
|
||||
"max": 0.1009521484375
|
||||
},
|
||||
"log_one_minus_alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": -7.0703125,
|
||||
"max": -0.004669189453125
|
||||
},
|
||||
"model_ema.decay": {
|
||||
"shape": [],
|
||||
"min": 1.0,
|
||||
"max": 1.0
|
||||
},
|
||||
"model_ema.num_updates": {
|
||||
"shape": [],
|
||||
"min": 219996,
|
||||
"max": 219996
|
||||
},
|
||||
"posterior_log_variance_clipped": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": -46.0625,
|
||||
"max": -4.421875
|
||||
},
|
||||
"posterior_mean_coef1": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.000827789306640625,
|
||||
"max": 1.0
|
||||
},
|
||||
"posterior_mean_coef2": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0,
|
||||
"max": 0.99560546875
|
||||
},
|
||||
"posterior_variance": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0,
|
||||
"max": 0.01200103759765625
|
||||
},
|
||||
"sqrt_alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0682373046875,
|
||||
"max": 0.99951171875
|
||||
},
|
||||
"sqrt_one_minus_alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0291595458984375,
|
||||
"max": 0.99755859375
|
||||
},
|
||||
"sqrt_recip_alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 1.0,
|
||||
"max": 14.6484375
|
||||
},
|
||||
"sqrt_recipm1_alphas_cumprod": {
|
||||
"shape": [
|
||||
1000
|
||||
],
|
||||
"min": 0.0291595458984375,
|
||||
"max": 14.6171875
|
||||
}
|
||||
},
|
||||
"diffusers": {}
|
||||
}
|
||||
4154
toolkit/keymaps/stable_diffusion_sdxl.json
Normal file
4154
toolkit/keymaps/stable_diffusion_sdxl.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_sdxl_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_sdxl_ldm_base.safetensors
Normal file
Binary file not shown.
35
toolkit/keymaps/stable_diffusion_sdxl_unmatched.json
Normal file
35
toolkit/keymaps/stable_diffusion_sdxl_unmatched.json
Normal file
@@ -0,0 +1,35 @@
|
||||
{
|
||||
"ldm": {
|
||||
"conditioner.embedders.0.transformer.text_model.embeddings.position_ids": {
|
||||
"shape": [
|
||||
1,
|
||||
77
|
||||
],
|
||||
"min": 0.0,
|
||||
"max": 76.0
|
||||
},
|
||||
"conditioner.embedders.1.model.logit_scale": {
|
||||
"shape": [],
|
||||
"min": 4.60546875,
|
||||
"max": 4.60546875
|
||||
},
|
||||
"conditioner.embedders.1.model.text_projection": {
|
||||
"shape": [
|
||||
1280,
|
||||
1280
|
||||
],
|
||||
"min": -0.15966796875,
|
||||
"max": 0.230712890625
|
||||
}
|
||||
},
|
||||
"diffusers": {
|
||||
"te1_text_projection.weight": {
|
||||
"shape": [
|
||||
1280,
|
||||
1280
|
||||
],
|
||||
"min": -0.15966796875,
|
||||
"max": 0.230712890625
|
||||
}
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user