Compare commits
321
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fb8ee98a1c | ||
|
|
7ff02a35ce | ||
|
|
21200b41b9 | ||
|
|
c21605bb95 | ||
|
|
4880e3f096 | ||
|
|
35e76fdef0 | ||
|
|
871e425f3e | ||
|
|
e6534f8d1d | ||
|
|
b5ad0b659e | ||
|
|
d6a205cc7b | ||
|
|
f21efd7840 | ||
|
|
29855d94ba | ||
|
|
006a3b9750 | ||
|
|
b166f7c274 | ||
|
|
f6cfd6d0bc | ||
|
|
6a4fa1ef8d | ||
|
|
3cd17086f3 | ||
|
|
75137a7539 | ||
|
|
ee91f8c0c7 | ||
|
|
598c4204be | ||
|
|
8da6c91af8 | ||
|
|
01523350e1 | ||
|
|
c891203365 | ||
|
|
85599dc725 | ||
|
|
af16d1edcc | ||
|
|
b3db4ff3bb | ||
|
|
2a0761ae75 | ||
|
|
d8ee9ee522 | ||
|
|
4f084d251c | ||
|
|
fb521e9dbd | ||
|
|
f937106f6c | ||
|
|
5f8fe5c4c8 | ||
|
|
2f5eaeba7a | ||
|
|
7d8c845765 | ||
|
|
340336b53d | ||
|
|
5abed1487f | ||
|
|
b5cf857b1f | ||
|
|
e1813c75d1 | ||
|
|
4fde1b7dd6 | ||
|
|
892cf61c52 | ||
|
|
fdd531d3c5 | ||
|
|
ef6635fafb | ||
|
|
faf864ce15 | ||
|
|
b02221b23c | ||
|
|
d8b3175b57 | ||
|
|
a80c934679 | ||
|
|
813e95592d | ||
|
|
41ddfbcca6 | ||
|
|
8cd935251d | ||
|
|
0ed83dc656 | ||
|
|
3ebbcf726b | ||
|
|
3883adb697 | ||
|
|
2033b13370 | ||
|
|
3dabbfd55c | ||
|
|
0d8f5ca929 | ||
|
|
128adb99ab | ||
|
|
46ce69a29a | ||
|
|
135959ba0e | ||
|
|
c845c76b15 | ||
|
|
f6f50e5004 | ||
|
|
e9b5ca0dcc | ||
|
|
eb144d5212 | ||
|
|
3f7ba82a25 | ||
|
|
6e32be4dfd | ||
|
|
3e17bed208 | ||
|
|
9102e8e5e5 | ||
|
|
03c6175007 | ||
|
|
13ee2ad6b5 | ||
|
|
e665169ce8 | ||
|
|
de99482825 | ||
|
|
0d0ad2a4d0 | ||
|
|
33b4ac7cc2 | ||
|
|
3f7298633d | ||
|
|
fec9cbfadd | ||
|
|
06b3b7b459 | ||
|
|
70a645ac79 | ||
|
|
fa76a90c39 | ||
|
|
8b3f9a930f | ||
|
|
cfde9e9180 | ||
|
|
4389ae75c8 | ||
|
|
d66863b4cb | ||
|
|
110ec0df30 | ||
|
|
b51727219e | ||
|
|
b8fe89e983 | ||
|
|
e47fcf429e | ||
|
|
fc63399317 | ||
|
|
0bd26033cd | ||
|
|
7d78a3d76d | ||
|
|
45e6f50c24 | ||
|
|
39ee9c9630 | ||
|
|
24fcc2e791 | ||
|
|
8da2f58c20 | ||
|
|
7d469d843c | ||
|
|
2e5912f629 | ||
|
|
37bcfe61c9 | ||
|
|
af6c4676d5 | ||
|
|
9ae9bc19d2 | ||
|
|
7218c7e8c2 | ||
|
|
c60f3646fa | ||
|
|
e1b05b1b99 | ||
|
|
0a72748da6 | ||
|
|
f919e3e7a1 | ||
|
|
d4923ab6d2 | ||
|
|
d613eff061 | ||
|
|
b48419247a | ||
|
|
4d9f8c2e67 | ||
|
|
34e4bc227c | ||
|
|
c85c45b688 | ||
|
|
d4dbb52f9b | ||
|
|
784bfbf14f | ||
|
|
327339ebba | ||
|
|
f2d70b64da | ||
|
|
41b7895bff | ||
|
|
5b4020d448 | ||
|
|
2633c06a87 | ||
|
|
a33a33295e | ||
|
|
3e44dff1da | ||
|
|
b2080da7c8 | ||
|
|
13733a4f19 | ||
|
|
ccaeab8e71 | ||
|
|
b568c0c283 | ||
|
|
c071905658 | ||
|
|
d1d8eb125e | ||
|
|
ea67904a9a | ||
|
|
909e1c0b10 | ||
|
|
d573407319 | ||
|
|
6f5253c1d4 | ||
|
|
f5a1de5ffe | ||
|
|
6470658e08 | ||
|
|
2af80788ca | ||
|
|
284b934567 | ||
|
|
25c9afe587 | ||
|
|
527f075d7a | ||
|
|
a5b7621732 | ||
|
|
6a6b5df282 | ||
|
|
b4891984aa | ||
|
|
9ac8d69bfc | ||
|
|
223969ed84 | ||
|
|
57d5b6b762 | ||
|
|
2a7c53a0c4 | ||
|
|
4f64882f11 | ||
|
|
1bd9e875fc | ||
|
|
bcaccf01c9 | ||
|
|
0bd726b6be | ||
|
|
cde4432588 | ||
|
|
aaf090c2e8 | ||
|
|
5181bbdff6 | ||
|
|
29d62cd705 | ||
|
|
9a985de529 | ||
|
|
fd0c8cde12 | ||
|
|
fcc569c0a6 | ||
|
|
7a70c18c05 | ||
|
|
a4a4d9aaa3 | ||
|
|
c4ea271482 | ||
|
|
66b34b312e | ||
|
|
0f9f5a7009 | ||
|
|
b8c668822b | ||
|
|
85b86f3667 | ||
|
|
8d58e312ed | ||
|
|
a4b8149ff3 | ||
|
|
4385a55a23 | ||
|
|
de0d8a8a41 | ||
|
|
b9a315bb98 | ||
|
|
c788afa996 | ||
|
|
e606340c56 | ||
|
|
08c7394677 | ||
|
|
54ba9492bc | ||
|
|
13ba5da163 | ||
|
|
430c8c62d6 | ||
|
|
9d63a9c1fd | ||
|
|
ff7b53ee63 | ||
|
|
744051b7b3 | ||
|
|
e19b73f3dd | ||
|
|
501f9991ab | ||
|
|
0e5867c4f8 | ||
|
|
8d970affc1 | ||
|
|
22737bb36c | ||
|
|
bda436a5b7 | ||
|
|
493310e51f | ||
|
|
1ede738580 | ||
|
|
eadb18c0d0 | ||
|
|
6aff99be49 | ||
|
|
221925c6ef | ||
|
|
2eeecae5d1 | ||
|
|
8ec0060ecd | ||
|
|
72eae9e1eb | ||
|
|
a474550c1b | ||
|
|
78011776f7 | ||
|
|
c01a99c17f | ||
|
|
ceb91f16e2 | ||
|
|
f31643c98a | ||
|
|
303eb29ece | ||
|
|
f33743c29d | ||
|
|
2fe846eba2 | ||
|
|
85748c11f0 | ||
|
|
cec51e15d8 | ||
|
|
ea470f5de7 | ||
|
|
868e0e1b2f | ||
|
|
8bcced05f1 | ||
|
|
12588e3124 | ||
|
|
24482e5009 | ||
|
|
d642537a33 | ||
|
|
5fdd0c1172 | ||
|
|
129ba6855c | ||
|
|
46390a0802 | ||
|
|
cba5273546 | ||
|
|
7970aab689 | ||
|
|
d9a4e8f361 | ||
|
|
ad952988db | ||
|
|
3870c8ebda | ||
|
|
820edaf20b | ||
|
|
8230058e65 | ||
|
|
71fd001f78 | ||
|
|
51cd996b3c | ||
|
|
4e9d5073a7 | ||
|
|
0c7c225909 | ||
|
|
129e5f756a | ||
|
|
3a6aa8cec3 | ||
|
|
13d57aaaed | ||
|
|
bf91ee3e83 | ||
|
|
036a899400 | ||
|
|
386e56a83a | ||
|
|
e517bd335b | ||
|
|
eef877e206 | ||
|
|
f593fdf5a7 | ||
|
|
7d46a193dc | ||
|
|
14bfba282a | ||
|
|
3a55b517f2 | ||
|
|
e43cb3891f | ||
|
|
da9781fd69 | ||
|
|
a840283f64 | ||
|
|
9198b2adab | ||
|
|
76c445f073 | ||
|
|
8d1d4560c8 | ||
|
|
9dc392a7b9 | ||
|
|
660aa3fa16 | ||
|
|
5fb1e6c39c | ||
|
|
161d3ad17f | ||
|
|
0161bc5935 | ||
|
|
c73f6462d2 | ||
|
|
e277232e62 | ||
|
|
d404d59952 | ||
|
|
ba28aedd15 | ||
|
|
f25f40a470 | ||
|
|
d28214c850 | ||
|
|
cd62af4049 | ||
|
|
8b20f2da23 | ||
|
|
059f7f2645 | ||
|
|
6007d09366 | ||
|
|
cff6fd00a4 | ||
|
|
abead52a2f | ||
|
|
eaaab69d8b | ||
|
|
671ff7b8dc | ||
|
|
12ec54e069 | ||
|
|
cd355d6e7e | ||
|
|
ff0c34e9da | ||
|
|
65d759dfeb | ||
|
|
67445d3963 | ||
|
|
4e97b176a1 | ||
|
|
1728715cc2 | ||
|
|
90feb0709c | ||
|
|
d56af4c5ed | ||
|
|
e3aa07ccfb | ||
|
|
3f90825f7d | ||
|
|
6091dd7a2c | ||
|
|
11fd00e5b0 | ||
|
|
4a3112c5be | ||
|
|
746279f20b | ||
|
|
e50844e79f | ||
|
|
e0de792a9d | ||
|
|
8ce5088064 | ||
|
|
be368b78bb | ||
|
|
b1ab40b15c | ||
|
|
da7dbd6000 | ||
|
|
b4d6a0dda8 | ||
|
|
cd9b6629bc | ||
|
|
6b8715663d | ||
|
|
e19713a4b8 | ||
|
|
7d8b7b2c63 | ||
|
|
d2b30450e4 | ||
|
|
975e385503 | ||
|
|
0a25118461 | ||
|
|
9d383569bc | ||
|
|
c2468dd46d | ||
|
|
3ac41fbe14 | ||
|
|
e7e3eecb48 | ||
|
|
61944f27e8 | ||
|
|
364dfc8157 | ||
|
|
99187c1f67 | ||
|
|
904a5e6237 | ||
|
|
28c13c21f8 | ||
|
|
67574ef1b8 | ||
|
|
68e61b6c1d | ||
|
|
9d4bc0ae2c | ||
|
|
fac2811fc6 | ||
|
|
fb57183706 | ||
|
|
63def6ade0 | ||
|
|
5e893857bb | ||
|
|
2fef34a4b9 | ||
|
|
3e5a6cbe35 | ||
|
|
28eff44081 | ||
|
|
f2764cd309 | ||
|
|
ccf010cdb3 | ||
|
|
d72005a9c0 | ||
|
|
8547000fdb | ||
|
|
21a4c42da0 | ||
|
|
500e974f5e | ||
|
|
4673ec5ef4 | ||
|
|
0c8c0edbfd | ||
|
|
59538e6995 | ||
|
|
dc6023b2f8 | ||
|
|
130db42b77 | ||
|
|
c185b0f0d6 | ||
|
|
8326386934 | ||
|
|
039c93ad70 | ||
|
|
c95978088c | ||
|
|
42c31241af | ||
|
|
09aa083e1b | ||
|
|
d4916bafc9 | ||
|
|
435d043873 | ||
|
|
269eeec1bc |
@@ -0,0 +1,32 @@
|
||||
# The .dockerignore file excludes files from the container build process.
|
||||
# https://docs.docker.com/engine/reference/builder/#dockerignore-file
|
||||
|
||||
# Exclude Git files
|
||||
.git
|
||||
.github
|
||||
.gitignore
|
||||
|
||||
# Exclude Python cache files
|
||||
__pycache__
|
||||
.mypy_cache
|
||||
.pytest_cache
|
||||
.ruff_cache
|
||||
|
||||
# exclude trained model rars:
|
||||
*.rar
|
||||
|
||||
# Dev folders:
|
||||
debug
|
||||
xander
|
||||
datasets
|
||||
rendered_images
|
||||
|
||||
# trained models:
|
||||
lora_models/*
|
||||
|
||||
# Ignore the entire models folder by default:
|
||||
models/*
|
||||
|
||||
### Include pipeline models: ###
|
||||
!models/juggernaut_reborn.safetensors
|
||||
!models/juggernaut_v6.safetensors
|
||||
+17
-8
@@ -1,15 +1,24 @@
|
||||
models
|
||||
lora_models
|
||||
remove
|
||||
|
||||
cache
|
||||
__pycache__
|
||||
.ipynb_checkpoints/
|
||||
models
|
||||
lora_models*
|
||||
eden_lora_training_runs/
|
||||
datasets
|
||||
|
||||
*.tar
|
||||
.env
|
||||
.cog
|
||||
xander*.sh
|
||||
.huggingface
|
||||
tests/
|
||||
trainer/
|
||||
train.py
|
||||
rendered_images*
|
||||
|
||||
gridsearch*
|
||||
aesthetic_score_best_model.pth
|
||||
|
||||
# experiment folders:
|
||||
conditioning_spaces/
|
||||
training_args_x_*.json
|
||||
xander_configs/
|
||||
debug/*
|
||||
!debug/*.py
|
||||
|
||||
|
||||
@@ -1,9 +1,42 @@
|
||||
# Trainer
|
||||
|
||||
Code for finetuning and training LoRa modules on top of Stable Diffusion.
|
||||
This trainer was developed by the [**Eden** team](https://eden.art/)
|
||||
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
|
||||
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
|
||||
|
||||
The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
|
||||
|
||||
<p align="center">
|
||||
<strong>Training images:</strong><br>
|
||||
<img src="assets/xander_training_images.jpg" alt="Image 1" style="width:80%;"/>
|
||||
</p>
|
||||
<p align="center">
|
||||
<strong>Generated imgs with trained LoRa:</strong><br>
|
||||
<img src="assets/xander_generated_images.jpg" alt="Image 2" style="width:80%;"/>
|
||||
</p>
|
||||
|
||||
|
||||
The trainer supports 3 default modes:
|
||||
- **style**: used for learning the aesthetic style of a collection of images.
|
||||
- **face**: used for learning a specific face (can be human, character, ...).
|
||||
- **object**: will learn a specific object or thing featured in the training images.
|
||||
|
||||
## Setup
|
||||
|
||||
Install all dependencies using
|
||||
|
||||
`pip install -r requirements.txt`
|
||||
|
||||
then you can simply run:
|
||||
|
||||
`python main.py train_configs/training_args.json`
|
||||
to start a training job.
|
||||
|
||||
Adjust the arguments inside `training_args.json` to setup a custom training job.
|
||||
|
||||
---
|
||||
|
||||
You can also run this through Replicate using cog (~docker image):
|
||||
1. Install Replicate 'cog':
|
||||
|
||||
```
|
||||
@@ -11,50 +44,62 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
|
||||
sudo chmod +x /usr/local/bin/cog
|
||||
```
|
||||
|
||||
2. Build the image with `sudo cog build`
|
||||
3. Run a training run with `sudo sh test_train.sh`
|
||||
2. Build the image with `cog build`
|
||||
3. Run a training run with `sh cog_test_train.sh`
|
||||
4. You can also go into the container with `cog run /bin/bash`
|
||||
|
||||
## Automatic Checkpoint Evaluation
|
||||
|
||||
This script uses CLIP img/txt similarity scores to evaluate how good the LoRa is vs how overfit.
|
||||
Download the aesthetic predictor model checkpoint first from google drive. This should give you a file named: `aesthetic_score_best_model.pth` (99.2 MB)
|
||||
|
||||
```bash
|
||||
gdown 1thEIlXVc8lkULVUBY9Ab45tsOERxkjxns
|
||||
```
|
||||
|
||||
Once the model is downloaded, you can run the eval script with the following CLI args:
|
||||
|
||||
- `output_folder`: this is where the outputs of the model get saved as jpeg files
|
||||
- `lora_path`: path to your LoRA checkpoint (make sure you edit `path_to_your_model_checkpoints` to point to the correct folder. It generally ends with something like `checkpoint-600` where `600` was the training step)
|
||||
- `output_json`: save all scores in this json file
|
||||
- `config_filename`: config file used for training
|
||||
|
||||
```bash
|
||||
python3 evaluate.py \
|
||||
--output_folder eval_images \
|
||||
--lora_path path_to_your_model_checkpoint \
|
||||
--output_json eval_results.json \
|
||||
--config_filename training_args.json
|
||||
```
|
||||
|
||||
|
||||
## TODO's
|
||||
|
||||
Code / Cleanup:
|
||||
- turn all/most of the args of the main() function in trainer_pti.py and the preprocess() function into a clean args_dict that makes it easy to add and distribute new parameters over the code and save these args to a .json file at the end.
|
||||
- Modularize the logic in train.py as much as possible, trying to minimize dev work that needs to happen when SD3 drops
|
||||
- make a clean train.py entrypoint that can be run as a normal python command (instead of having to use cog)
|
||||
- make it so the textual_inversion optimizer only optimizes the actual trained token embeddings instead of all of them + resetting later
|
||||
- test if the trained concepts with peft are compatible with ComfyUI / AUTO1111
|
||||
Bugs:
|
||||
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
|
||||
|
||||
Algo:
|
||||
- Add aspect_ratio bucketing into the dataloader so we can train on non-square images (take this from https://github.com/kohya-ss/sd-scripts)
|
||||
- Improve some of the chatgpt functionality:
|
||||
- separate the "gpt_description" / "gpt_segmentation" prompt calls and make them run on a subset of prompts in case there's a lot of imgs / prompts (possibly use img_grids for some gpt4-v calls)
|
||||
- currently some sub-optimal stuff can happen in preprocess() when there's less than 3 or more than 45 imgs, try to improve this
|
||||
- Test if timesteps = torch.randint() can be improved: look at sdxl training code! (see https://github.com/huggingface/diffusers/blob/main/examples/advanced_diffusion_training/train_dreambooth_lora_sdxl_advanced.py#L1263, https://arxiv.org/pdf/2206.00364.pdf)
|
||||
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
|
||||
- test if textual inversion training can also happen with prodigy_optimizer
|
||||
- the random initialization of the token embeddings has a relatively large impact on the final outcome, there are prob ways to reduce
|
||||
this random variance, eg CLIP_similarity pretraining.
|
||||
- Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM
|
||||
|
||||
Bugfixing:
|
||||
see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037
|
||||
- check how the pipe() objects work under the hood in HF diffusers library, is there a difference w how the unet is called in the training loop / which args it gets?
|
||||
- Try to find out why the diffusers training script works for sd15 and ours doesnt:
|
||||
See here: https://huggingface.co/blog/sdxl_lora_advanced_script
|
||||
and here: https://github.com/huggingface/diffusers/tree/main/examples/advanced_diffusion_training
|
||||
- figure out how to adaptively set lora_scale at inference time using peft + diffusers? (https://github.com/huggingface/peft/blob/main/src/peft/tuners/lora/layer.py#L240)
|
||||
- improve data augmentation, eg by adding outpainted, smaller versions of faces / objects
|
||||
|
||||
Small, minor tweaks:
|
||||
- preprocess.py: the imgs are first auto-captioned and then cropped, this is not ideal, swap this around!
|
||||
|
||||
Bigger improvements:
|
||||
- add stronger token regularization (eg CelebBasis spanning basis):
|
||||
- Add multi-token training
|
||||
- implement perfusion: https://research.nvidia.com/labs/par/Perfusion/
|
||||
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
|
||||
- implement prompt-aligned: https://prompt-aligned.github.io/
|
||||
- make compatible with ziplora: https://ziplora.github.io/
|
||||
|
||||
|
||||
|
||||
Tuning Experiments once code is fully ready:
|
||||
|
||||
Tuning Experiments:
|
||||
- try-out conditioning noise injection during training to increase robustness
|
||||
- re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam)
|
||||
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
|
||||
but it then recovers. Can we avoid this collapse? Is the learning rate too high?
|
||||
- gradient_accumulation
|
||||
- offset noise
|
||||
- AB test Dora vs Lora
|
||||
- sweep n_trainable_tokens to inject
|
||||
|
||||
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
import os
|
||||
import sys
|
||||
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
|
||||
|
||||
from node import Eden_LoRa_trainer
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Eden_LoRa_trainer": Eden_LoRa_trainer,
|
||||
}
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 915 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 342 KiB |
@@ -3,36 +3,14 @@
|
||||
|
||||
build:
|
||||
gpu: true
|
||||
cuda: "11.8"
|
||||
python_version: "3.9"
|
||||
cuda: "12.1"
|
||||
python_version: "3.11"
|
||||
system_packages:
|
||||
- "libgl1-mesa-glx"
|
||||
- "ffmpeg"
|
||||
- "libsm6"
|
||||
- "libxext6"
|
||||
python_packages:
|
||||
- "scipy"
|
||||
- "diffusers==0.25.1"
|
||||
- "peft==0.9.0"
|
||||
- "torch==2.0.1"
|
||||
- "transformers==4.31.0"
|
||||
- "invisible-watermark==0.2.0"
|
||||
- "accelerate==0.21.0"
|
||||
- "pandas==2.0.3"
|
||||
- "torchvision==0.15.2"
|
||||
- "numpy==1.25.1"
|
||||
- "pandas==2.0.3"
|
||||
- "fire==0.5.0"
|
||||
- "opencv-python>=4.1.0.25"
|
||||
- "mediapipe==0.10.2"
|
||||
- "openai==1.2.4"
|
||||
- python-dotenv
|
||||
- prodigyopt
|
||||
- omegaconf
|
||||
|
||||
python_requirements: requirements.txt
|
||||
run:
|
||||
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/download/v0.0.1/pget" && chmod +x /usr/local/bin/pget
|
||||
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task
|
||||
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
|
||||
|
||||
predict: "predict.py:Predictor"
|
||||
image: "r8.im/abraham-ai/sdxl-lora-trainer"
|
||||
image: "r8.im/edenartlab/sdxl-lora-trainer"
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
# Set GPU ID to run these jobs on:
|
||||
GPU_ID="device=3"
|
||||
|
||||
cog predict --gpus $GPU_ID \
|
||||
-i name="xander_sdxl_cog" \
|
||||
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip" \
|
||||
-i concept_mode="face" \
|
||||
-i sd_model_version="sdxl" \
|
||||
-i max_train_steps="360" \
|
||||
-i caption_model="blip" \
|
||||
-i debug="False" \
|
||||
-i seed="0"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,533 @@
|
||||
{
|
||||
"last_node_id": 12,
|
||||
"last_link_id": 23,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 7,
|
||||
"type": "CLIPTextEncode",
|
||||
"pos": [
|
||||
413,
|
||||
389
|
||||
],
|
||||
"size": {
|
||||
"0": 425.27801513671875,
|
||||
"1": 180.6060791015625
|
||||
},
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 16
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"links": [
|
||||
6
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "CLIPTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"text, watermark"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VAEDecode",
|
||||
"pos": [
|
||||
1209,
|
||||
188
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
9
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VAEDecode"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 12,
|
||||
"type": "Reroute",
|
||||
"pos": [
|
||||
220,
|
||||
14
|
||||
],
|
||||
"size": [
|
||||
75,
|
||||
26
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "",
|
||||
"type": "*",
|
||||
"link": 22
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "",
|
||||
"type": "MODEL",
|
||||
"links": [
|
||||
19
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"showOutputText": false,
|
||||
"horizontal": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "KSampler",
|
||||
"pos": [
|
||||
863,
|
||||
186
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 262
|
||||
},
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "MODEL",
|
||||
"link": 19
|
||||
},
|
||||
{
|
||||
"name": "positive",
|
||||
"type": "CONDITIONING",
|
||||
"link": 4
|
||||
},
|
||||
{
|
||||
"name": "negative",
|
||||
"type": "CONDITIONING",
|
||||
"link": 6
|
||||
},
|
||||
{
|
||||
"name": "latent_image",
|
||||
"type": "LATENT",
|
||||
"link": 2
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"links": [
|
||||
7
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "KSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
1,
|
||||
"fixed",
|
||||
25,
|
||||
8,
|
||||
"euler",
|
||||
"normal",
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "EmptyLatentImage",
|
||||
"pos": [
|
||||
473,
|
||||
609
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 106
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "LATENT",
|
||||
"type": "LATENT",
|
||||
"links": [
|
||||
2
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "EmptyLatentImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
768,
|
||||
768,
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "CheckpointLoaderSimple",
|
||||
"pos": [
|
||||
-467,
|
||||
120
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 98
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"links": [
|
||||
10
|
||||
],
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"links": [
|
||||
12
|
||||
],
|
||||
"slot_index": 1
|
||||
},
|
||||
{
|
||||
"name": "VAE",
|
||||
"type": "VAE",
|
||||
"links": [
|
||||
8
|
||||
],
|
||||
"slot_index": 2
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "CheckpointLoaderSimple"
|
||||
},
|
||||
"widgets_values": [
|
||||
"juggernaut_reborn.safetensors"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "CLIPTextEncode",
|
||||
"pos": [
|
||||
415,
|
||||
186
|
||||
],
|
||||
"size": {
|
||||
"0": 422.84503173828125,
|
||||
"1": 164.31304931640625
|
||||
},
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 15
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "CONDITIONING",
|
||||
"type": "CONDITIONING",
|
||||
"links": [
|
||||
4
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "CLIPTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"a photo of embedding:xander_sd15_embedding on the beach "
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "Reroute",
|
||||
"pos": [
|
||||
220,
|
||||
49
|
||||
],
|
||||
"size": [
|
||||
75,
|
||||
26
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "",
|
||||
"type": "*",
|
||||
"link": 23
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "",
|
||||
"type": "CLIP",
|
||||
"links": [
|
||||
15,
|
||||
16
|
||||
],
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"showOutputText": false,
|
||||
"horizontal": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 9,
|
||||
"type": "SaveImage",
|
||||
"pos": [
|
||||
347,
|
||||
-270
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 270
|
||||
},
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 9
|
||||
}
|
||||
],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"ComfyUI"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 10,
|
||||
"type": "LoraLoader",
|
||||
"pos": [
|
||||
-72,
|
||||
-229
|
||||
],
|
||||
"size": {
|
||||
"0": 254.95774841308594,
|
||||
"1": 127.86701202392578
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "MODEL",
|
||||
"link": 10
|
||||
},
|
||||
{
|
||||
"name": "clip",
|
||||
"type": "CLIP",
|
||||
"link": 12
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MODEL",
|
||||
"type": "MODEL",
|
||||
"links": [
|
||||
22
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "CLIP",
|
||||
"type": "CLIP",
|
||||
"links": [
|
||||
23
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoraLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"xander_sd15_lora.safetensors",
|
||||
0.6,
|
||||
0.6
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
2,
|
||||
5,
|
||||
0,
|
||||
3,
|
||||
3,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
4,
|
||||
6,
|
||||
0,
|
||||
3,
|
||||
1,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
6,
|
||||
7,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"CONDITIONING"
|
||||
],
|
||||
[
|
||||
7,
|
||||
3,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
8,
|
||||
4,
|
||||
2,
|
||||
8,
|
||||
1,
|
||||
"VAE"
|
||||
],
|
||||
[
|
||||
9,
|
||||
8,
|
||||
0,
|
||||
9,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
10,
|
||||
4,
|
||||
0,
|
||||
10,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
12,
|
||||
4,
|
||||
1,
|
||||
10,
|
||||
1,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
15,
|
||||
11,
|
||||
0,
|
||||
6,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
16,
|
||||
11,
|
||||
0,
|
||||
7,
|
||||
0,
|
||||
"CLIP"
|
||||
],
|
||||
[
|
||||
19,
|
||||
12,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"MODEL"
|
||||
],
|
||||
[
|
||||
22,
|
||||
10,
|
||||
0,
|
||||
12,
|
||||
0,
|
||||
"*"
|
||||
],
|
||||
[
|
||||
23,
|
||||
10,
|
||||
1,
|
||||
11,
|
||||
0,
|
||||
"*"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8264462809917354,
|
||||
"offset": {
|
||||
"0": 513.8734070325743,
|
||||
"1": 351.4824273966635
|
||||
}
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -1,805 +0,0 @@
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import random
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import gc
|
||||
import PIL
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from PIL import Image
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import save_file
|
||||
from torch.utils.data import Dataset
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
def plot_torch_hist(parameters, epoch, checkpoint_dir, name, bins=100, min_val=-1, max_val=1, ymax_f = 0.75):
|
||||
# Flatten and concatenate all parameters into a single tensor
|
||||
all_params = torch.cat([p.data.view(-1) for p in parameters])
|
||||
|
||||
# Convert to CPU for plotting
|
||||
all_params_cpu = all_params.cpu().float().numpy()
|
||||
|
||||
# Plot histogram
|
||||
plt.figure()
|
||||
plt.hist(all_params_cpu, bins=bins, density=False)
|
||||
plt.ylim(0, ymax_f * len(all_params_cpu))
|
||||
plt.xlim(min_val, max_val)
|
||||
plt.xlabel('Weight Value')
|
||||
plt.ylabel('Count')
|
||||
plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f})')
|
||||
plt.savefig(f"{checkpoint_dir}/{name}_histogram_{epoch:04d}.png")
|
||||
plt.close()
|
||||
|
||||
# plot the learning rates:
|
||||
def plot_lrs(lora_lrs, ti_lrs, save_path='learning_rates.png'):
|
||||
plt.figure()
|
||||
plt.plot(range(len(lora_lrs)), lora_lrs, label='LoRA LR')
|
||||
plt.plot(range(len(lora_lrs)), ti_lrs, label='TI LR')
|
||||
plt.yscale('log') # Set y-axis to log scale
|
||||
plt.ylim(1e-6, 3e-3)
|
||||
plt.xlabel('Step')
|
||||
plt.ylabel('Learning Rate')
|
||||
plt.title('Learning Rate Curves')
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
from scipy.signal import savgol_filter
|
||||
def plot_loss(losses, save_path='losses.png', window_length=31, polyorder=3):
|
||||
if len(losses) < window_length:
|
||||
return
|
||||
|
||||
smoothed_losses = savgol_filter(losses, window_length, polyorder)
|
||||
|
||||
plt.figure()
|
||||
plt.plot(losses, label='Actual Losses')
|
||||
plt.plot(smoothed_losses, label='Smoothed Losses', color='red')
|
||||
# plt.yscale('log') # Uncomment if log scale is desired
|
||||
plt.xlabel('Step')
|
||||
plt.ylabel('Training Loss')
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
|
||||
def prepare_image(
|
||||
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
|
||||
) -> torch.Tensor:
|
||||
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
|
||||
arr = np.array(pil_image.convert("RGB"))
|
||||
arr = arr.astype(np.float32) / 127.5 - 1
|
||||
arr = np.transpose(arr, [2, 0, 1])
|
||||
image = torch.from_numpy(arr).unsqueeze(0)
|
||||
return image
|
||||
|
||||
|
||||
def prepare_mask(
|
||||
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
|
||||
) -> torch.Tensor:
|
||||
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
|
||||
arr = np.array(pil_image.convert("L"))
|
||||
arr = arr.astype(np.float32) / 255.0
|
||||
arr = np.expand_dims(arr, 0)
|
||||
image = torch.from_numpy(arr).unsqueeze(0)
|
||||
return image
|
||||
|
||||
|
||||
class PreprocessedDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
csv_path: str,
|
||||
tokenizer_1,
|
||||
tokenizer_2,
|
||||
vae_encoder,
|
||||
text_encoder_1=None,
|
||||
text_encoder_2=None,
|
||||
do_cache: bool = False,
|
||||
size: int = 512,
|
||||
text_dropout: float = 0.0,
|
||||
scale_vae_latents: bool = True,
|
||||
substitute_caption_map: Dict[str, str] = {},
|
||||
):
|
||||
super().__init__()
|
||||
self.data = pd.read_csv(csv_path)
|
||||
self.csv_path = csv_path
|
||||
|
||||
self.caption = self.data["caption"]
|
||||
# make it lowercase
|
||||
self.caption = self.caption.str.lower()
|
||||
for key, value in substitute_caption_map.items():
|
||||
self.caption = self.caption.str.replace(key.lower(), value)
|
||||
|
||||
self.image_path = self.data["image_path"]
|
||||
|
||||
if "mask_path" not in self.data.columns:
|
||||
self.mask_path = None
|
||||
else:
|
||||
self.mask_path = self.data["mask_path"]
|
||||
|
||||
if text_encoder_1 is None:
|
||||
self.return_text_embeddings = False
|
||||
else:
|
||||
self.text_encoder_1 = text_encoder_1
|
||||
self.text_encoder_2 = text_encoder_2
|
||||
self.return_text_embeddings = True
|
||||
assert (
|
||||
NotImplementedError
|
||||
), "Preprocessing Text Encoder is not implemented yet"
|
||||
|
||||
self.tokenizer_1 = tokenizer_1
|
||||
self.tokenizer_2 = tokenizer_2
|
||||
|
||||
self.vae_encoder = vae_encoder
|
||||
self.scale_vae_latents = scale_vae_latents
|
||||
self.text_dropout = text_dropout
|
||||
self.size = size
|
||||
|
||||
if do_cache:
|
||||
self.vae_latents = []
|
||||
self.tokens_tuple = []
|
||||
self.masks = []
|
||||
|
||||
self.do_cache = True
|
||||
|
||||
print("Captions to train on: ")
|
||||
for idx in range(len(self.data)):
|
||||
token, vae_latent, mask = self._process(idx)
|
||||
self.vae_latents.append(vae_latent)
|
||||
self.tokens_tuple.append(token)
|
||||
self.masks.append(mask)
|
||||
|
||||
print(f"Cached latents and masks for {len(self.vae_latents)} images.")
|
||||
|
||||
del self.vae_encoder
|
||||
|
||||
else:
|
||||
self.do_cache = False
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.data)
|
||||
|
||||
@torch.no_grad()
|
||||
def _process(
|
||||
self, idx: int
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
|
||||
image_path = self.image_path[idx]
|
||||
image_path = os.path.join(os.path.dirname(self.csv_path), image_path)
|
||||
|
||||
image = PIL.Image.open(image_path).convert("RGB")
|
||||
|
||||
image = prepare_image(image, self.size, self.size).to(
|
||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
caption = self.caption[idx]
|
||||
print(caption)
|
||||
|
||||
# tokenizer_1
|
||||
ti1 = self.tokenizer_1(
|
||||
caption,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
).input_ids.squeeze()
|
||||
|
||||
if self.tokenizer_2 is None:
|
||||
ti2 = None
|
||||
else:
|
||||
ti2 = self.tokenizer_2(
|
||||
caption,
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
).input_ids.squeeze()
|
||||
|
||||
vae_latent = self.vae_encoder.encode(image).latent_dist.sample()
|
||||
|
||||
if self.scale_vae_latents:
|
||||
vae_latent = vae_latent * self.vae_encoder.config.scaling_factor
|
||||
|
||||
if self.mask_path is None:
|
||||
mask = torch.ones_like(
|
||||
vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
else:
|
||||
mask_path = self.mask_path[idx]
|
||||
mask_path = os.path.join(os.path.dirname(self.csv_path), mask_path)
|
||||
|
||||
mask = PIL.Image.open(mask_path)
|
||||
|
||||
mask = prepare_mask(mask, self.size, self.size).to(
|
||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
mask_dtype = mask.dtype
|
||||
mask = mask.float()
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask, size=(vae_latent.shape[-2], vae_latent.shape[-1]), mode="nearest"
|
||||
)
|
||||
mask = mask.to(dtype=mask_dtype)
|
||||
mask = mask.repeat(1, vae_latent.shape[1], 1, 1)
|
||||
|
||||
assert len(mask.shape) == 4 and len(vae_latent.shape) == 4
|
||||
|
||||
if ti2 is None: # sd15
|
||||
return ti1, vae_latent.squeeze(), mask.squeeze()
|
||||
else: # sdxl
|
||||
return (ti1, ti2), vae_latent.squeeze(), mask.squeeze()
|
||||
|
||||
def atidx(
|
||||
self, idx: int
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
|
||||
if self.do_cache:
|
||||
return self.tokens_tuple[idx], self.vae_latents[idx], self.masks[idx]
|
||||
else:
|
||||
return self._process(idx)
|
||||
|
||||
def __getitem__(
|
||||
self, idx: int
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
|
||||
token, vae_latent, mask = self.atidx(idx)
|
||||
return token, vae_latent, mask
|
||||
|
||||
|
||||
def import_model_class_from_model_name_or_path(
|
||||
pretrained_model_name_or_path: str, revision: str, subfolder: str = "text_encoder"
|
||||
):
|
||||
text_encoder_config = PretrainedConfig.from_pretrained(
|
||||
pretrained_model_name_or_path, subfolder=subfolder, revision=revision
|
||||
)
|
||||
model_class = text_encoder_config.architectures[0]
|
||||
|
||||
if model_class == "CLIPTextModel":
|
||||
from transformers import CLIPTextModel
|
||||
print("Importing CLIPTextModel")
|
||||
return CLIPTextModel
|
||||
elif model_class == "CLIPTextModelWithProjection":
|
||||
from transformers import CLIPTextModelWithProjection
|
||||
print("Importing CLIPTextModelWithProjection")
|
||||
return CLIPTextModelWithProjection
|
||||
else:
|
||||
raise ValueError(f"{model_class} is not supported.")
|
||||
|
||||
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
|
||||
if not isinstance(pretrained_model, dict) or 'path' not in pretrained_model or 'version' not in pretrained_model:
|
||||
raise ValueError("pretrained_model must be a dict with 'path' and 'version' keys")
|
||||
|
||||
print(f"Loading model weights from {pretrained_model['path']} as dtype: {weight_dtype}...")
|
||||
|
||||
if pretrained_model['version'] == "sd15":
|
||||
pipe = StableDiffusionPipeline.from_single_file(
|
||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||
else:
|
||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||
|
||||
pipe = pipe.to(device, dtype=weight_dtype)
|
||||
|
||||
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
|
||||
vae = pipe.vae
|
||||
unet = pipe.unet
|
||||
tokenizer_one = pipe.tokenizer
|
||||
text_encoder_one = pipe.text_encoder
|
||||
|
||||
vae.requires_grad_(False)
|
||||
text_encoder_one.requires_grad_(False)
|
||||
|
||||
text_encoder_one.to(device, dtype=weight_dtype)
|
||||
unet.to(device, dtype=weight_dtype)
|
||||
if keep_vae_float32:
|
||||
vae.to(device, dtype=torch.float32)
|
||||
else:
|
||||
vae.to(device, dtype=weight_dtype)
|
||||
if weight_dtype != torch.float32:
|
||||
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but not for training!!")
|
||||
|
||||
tokenizer_two = text_encoder_two = None
|
||||
if pretrained_model['version'] == "sdxl":
|
||||
tokenizer_two = pipe.tokenizer_2
|
||||
text_encoder_two = pipe.text_encoder_2
|
||||
text_encoder_two.requires_grad_(False)
|
||||
text_encoder_two.to(device, dtype=weight_dtype)
|
||||
|
||||
return (
|
||||
pipe,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
noise_scheduler,
|
||||
text_encoder_one,
|
||||
text_encoder_two,
|
||||
vae,
|
||||
unet,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class TokenEmbeddingsHandler:
|
||||
def __init__(self, text_encoders, tokenizers):
|
||||
self.text_encoders = text_encoders
|
||||
self.tokenizers = tokenizers
|
||||
|
||||
self.train_ids: Optional[torch.Tensor] = None
|
||||
self.inserting_toks: Optional[List[str]] = None
|
||||
self.embeddings_settings = {}
|
||||
|
||||
|
||||
def get_trainable_embeddings(self):
|
||||
|
||||
trainable_embeddings = []
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
trainable_embeddings.append(text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids])
|
||||
|
||||
return trainable_embeddings
|
||||
|
||||
def find_nearest_tokens(self, query_embedding, tokenizer, text_encoder, idx, distance_metric, top_k = 5):
|
||||
# given a query embedding, compute the distance to all embeddings in the text encoder
|
||||
# and return the top_k closest tokens
|
||||
|
||||
assert distance_metric in ["l2", "cosine"], "distance_metric should be either 'l2' or 'cosine'"
|
||||
|
||||
# get all non-optimized embeddings:
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[index_no_updates]
|
||||
|
||||
# compute the distance between the query embedding and all embeddings:
|
||||
if distance_metric == "l2":
|
||||
diff = (embeddings - query_embedding.unsqueeze(0))**2
|
||||
distances = diff.sum(-1)
|
||||
distances, indices = torch.topk(distances, top_k, dim=0, largest=False)
|
||||
elif distance_metric == "cosine":
|
||||
distances = F.cosine_similarity(embeddings, query_embedding.unsqueeze(0), dim=-1)
|
||||
distances, indices = torch.topk(distances, top_k, dim=0, largest=True)
|
||||
|
||||
nearest_tokens = tokenizer.convert_ids_to_tokens(indices)
|
||||
return nearest_tokens, distances
|
||||
|
||||
|
||||
def print_token_info(self, distance_metric = "cosine"):
|
||||
print(f"----------- Closest tokens (distance_metric = {distance_metric}) --------------")
|
||||
current_token_embeddings = self.get_trainable_embeddings()
|
||||
idx = 0
|
||||
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if text_encoder is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
query_embeddings = current_token_embeddings[idx]
|
||||
|
||||
for token_id, query_embedding in enumerate(query_embeddings):
|
||||
nearest_tokens, distances = self.find_nearest_tokens(query_embedding, tokenizer, text_encoder, idx, distance_metric)
|
||||
|
||||
# print the results:
|
||||
print(f"txt-encoder {idx}, token {token_id}: :")
|
||||
for i, (token, dist) in enumerate(zip(nearest_tokens, distances)):
|
||||
print(f"---> {distance_metric} of {dist:.4f}: {token}")
|
||||
|
||||
idx += 1
|
||||
|
||||
def get_start_embedding(self, text_encoder, tokenizer, example_tokens, unk_token_id = 49407, verbose = False, desired_std_multiplier = 0.0):
|
||||
print('-----------------------------------------------')
|
||||
# do some cleanup:
|
||||
example_tokens = [tok.lower() for tok in example_tokens]
|
||||
example_tokens = list(set(example_tokens))
|
||||
|
||||
starting_ids = tokenizer.convert_tokens_to_ids(example_tokens)
|
||||
|
||||
# filter out any tokens that are mapped to unk_token_id:
|
||||
example_tokens = [tok for tok, tok_id in zip(example_tokens, starting_ids) if tok_id != unk_token_id]
|
||||
starting_ids = [tok_id for tok_id in starting_ids if tok_id != unk_token_id]
|
||||
|
||||
if verbose:
|
||||
print("Token mapping:")
|
||||
for i, token in enumerate(example_tokens):
|
||||
print(f"{token} -> {starting_ids[i]}")
|
||||
|
||||
embeddings, stds = [], []
|
||||
for i, token_index in enumerate(starting_ids):
|
||||
embedding = text_encoder.text_model.embeddings.token_embedding.weight.data[token_index].clone()
|
||||
embeddings.append(embedding)
|
||||
stds.append(embedding.std())
|
||||
#print(f"token: {example_tokens[i]}, embedding-std: {embedding.std():.4f}, embedding-mean: {embedding.mean():.4f}")
|
||||
|
||||
embeddings = torch.stack(embeddings)
|
||||
#print(f"Embeddings: {embeddings.shape}, std: {embeddings.std():.4f}, mean: {embeddings.mean():.4f}")
|
||||
|
||||
if verbose:
|
||||
# Compute the squared difference
|
||||
squared_diff = (embeddings.unsqueeze(1) - embeddings.unsqueeze(0)) ** 2
|
||||
squared_l2_dist = squared_diff.sum(-1)
|
||||
l2_distance_matrix = torch.sqrt(squared_l2_dist)
|
||||
|
||||
print("Pairwise L2 Distance Matrix:")
|
||||
print(" \t" + "\t".join(example_tokens))
|
||||
for i, row in enumerate(l2_distance_matrix):
|
||||
print(f"{example_tokens[i]}\t" + "\t".join(f"{dist:.4f}" for dist in row))
|
||||
|
||||
|
||||
# We're working in cosine-similarity space
|
||||
# So first, renormalize the embeddings to have norm 1
|
||||
embedding_norms = torch.norm(embeddings, dim=-1, keepdim=True)
|
||||
embeddings = embeddings / embedding_norms
|
||||
|
||||
print(f"embedding norms pre normalization:")
|
||||
print(embedding_norms)
|
||||
print(f"embedding norms post normalization:")
|
||||
print(torch.norm(embeddings, dim=-1, keepdim=True))
|
||||
|
||||
print(f"Using {len(embeddings)} embeddings to compute initial embedding...")
|
||||
init_embedding = embeddings.mean(dim=0)
|
||||
# normalize the init_embedding to have norm 1:
|
||||
init_embedding = init_embedding / torch.norm(init_embedding)
|
||||
|
||||
# rescale the init_embedding to have the same std as the average of the embeddings:
|
||||
init_embedding = init_embedding * embedding_norms.mean()
|
||||
|
||||
print(f"init_embedding norm: {torch.norm(init_embedding):.4f}, std: {init_embedding.std():.4f}, mean: {init_embedding.mean():.4f}")
|
||||
|
||||
if (desired_std_multiplier is not None) and desired_std_multiplier > 0:
|
||||
avg_std = torch.stack(stds).mean()
|
||||
current_std = init_embedding.std()
|
||||
scale_factor = desired_std_multiplier * avg_std / current_std
|
||||
init_embedding = init_embedding * scale_factor
|
||||
print(f"Scaled Mean Embedding: std: {init_embedding.std():.4f}, mean: {init_embedding.mean():.4f}")
|
||||
|
||||
return init_embedding
|
||||
|
||||
def plot_token_embeddings(self, example_tokens, output_folder = ".", x_range = [-0.05, 0.05]):
|
||||
print(f"Plotting embeddings for tokens: {example_tokens}")
|
||||
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
token_ids = tokenizer.convert_tokens_to_ids(example_tokens)
|
||||
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_ids].clone()
|
||||
|
||||
# plot the embeddings histogram:
|
||||
for token_name, embedding in zip(example_tokens, embeddings):
|
||||
plot_torch_hist(embedding, 0, output_folder, f"tok_{token_name}_{idx}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
|
||||
|
||||
idx += 1
|
||||
|
||||
def initialize_new_tokens(self,
|
||||
inserting_toks: List[str],
|
||||
starting_toks: Optional[List[str]] = None,
|
||||
seed: int = 0,
|
||||
):
|
||||
|
||||
print("Initializing new tokens...")
|
||||
print(inserting_toks)
|
||||
torch.manual_seed(seed)
|
||||
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
assert isinstance(
|
||||
inserting_toks, list
|
||||
), "inserting_toks should be a list of strings."
|
||||
assert all(
|
||||
isinstance(tok, str) for tok in inserting_toks
|
||||
), "All elements in inserting_toks should be strings."
|
||||
|
||||
self.inserting_toks = inserting_toks
|
||||
|
||||
print(f"Inserting new tokens into tokenizer-{idx}:")
|
||||
print(self.inserting_toks)
|
||||
|
||||
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
|
||||
tokenizer.add_special_tokens(special_tokens_dict)
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
|
||||
|
||||
# random initialization of new tokens
|
||||
std_token_embedding = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data.std() #(axis=1).mean()
|
||||
)
|
||||
std_token_mean = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data.mean() #(axis=1).mean()
|
||||
)
|
||||
|
||||
print(f"Text encoder {idx} token_embedding_std: {std_token_embedding}")
|
||||
|
||||
if starting_toks is not None:
|
||||
assert isinstance(
|
||||
starting_toks, list
|
||||
), "starting_toks should be a list of strings."
|
||||
assert all(
|
||||
isinstance(tok, str) for tok in starting_toks
|
||||
), "All elements in starting_toks should be strings."
|
||||
assert len(starting_toks) == len(self.inserting_toks), "starting_toks should have the same length as inserting_toks"
|
||||
self.starting_ids = tokenizer.convert_tokens_to_ids(starting_toks)
|
||||
|
||||
print(f"Copying embeddings from starting tokens {starting_toks} to new tokens {self.inserting_toks}")
|
||||
print(f"Starting ids: {self.starting_ids}")
|
||||
|
||||
# copy the embeddings of the starting tokens to the new tokens
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
|
||||
|
||||
else:
|
||||
|
||||
if 1: # random initialization:
|
||||
init_embeddings = (torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype) * std_token_embedding * 1.0)
|
||||
else:
|
||||
# Test code to initialize the new tokens with some specific tokens
|
||||
first_tokens = [
|
||||
"Sophia",
|
||||
"Liam",
|
||||
"Ethan",
|
||||
"Lucas",
|
||||
"Olivia",
|
||||
"Noah",
|
||||
"John",
|
||||
"David",
|
||||
"James",
|
||||
"Robert",
|
||||
"Michael",
|
||||
"William",
|
||||
]
|
||||
|
||||
second_tokens = [
|
||||
"Smith",
|
||||
"Johnson",
|
||||
"Williams",
|
||||
"Brown",
|
||||
"Jones",
|
||||
"Garcia",
|
||||
"Miller",
|
||||
"Davis",
|
||||
"Rodriguez",
|
||||
"Carter",
|
||||
"Trump",
|
||||
"Clinton",
|
||||
"Wilson",
|
||||
"Harris",
|
||||
"Lewis",
|
||||
"Scott"
|
||||
]
|
||||
|
||||
self.anchor_embedding_one = self.get_start_embedding(text_encoder, tokenizer, first_tokens)
|
||||
self.anchor_embedding_two = self.get_start_embedding(text_encoder, tokenizer, second_tokens)
|
||||
self.anchor_embedding_three = self.get_start_embedding(text_encoder, tokenizer, first_tokens)
|
||||
self.anchor_embedding_four = self.get_start_embedding(text_encoder, tokenizer, second_tokens)
|
||||
|
||||
init_embeddings = torch.stack([self.anchor_embedding_one, self.anchor_embedding_two, self.anchor_embedding_three, self.anchor_embedding_four])
|
||||
|
||||
print(f"init_embedding std: {init_embeddings.std():.4f}, avg-std: {std_token_embedding:.4f}")
|
||||
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
|
||||
|
||||
self.embeddings_settings[
|
||||
f"original_embeddings_{idx}"
|
||||
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
|
||||
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
|
||||
|
||||
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
inu[self.train_ids] = False
|
||||
|
||||
self.embeddings_settings[f"index_no_updates_{idx}"] = inu
|
||||
|
||||
idx += 1
|
||||
|
||||
def pre_optimize_token_embeddings(self, train_dataset, epochs=10):
|
||||
|
||||
### THIS FUNCTION IS NOT DONE YET
|
||||
### Idea here is to use CLIP-similarity between imgs and prompts to pre-optimize the embeddings
|
||||
|
||||
for idx in range(len(train_dataset)):
|
||||
(tok1, tok2), vae_latent, mask = train_dataset[idx]
|
||||
image_path = train_dataset.image_path[idx]
|
||||
image_path = os.path.join(os.path.dirname(train_dataset.csv_path), image_path)
|
||||
image = PIL.Image.open(image_path).convert("RGB")
|
||||
|
||||
print(f"---> Loaded sample {idx}:")
|
||||
print("Tokens:")
|
||||
print(tok1.shape)
|
||||
print(tok2.shape)
|
||||
print("Image:")
|
||||
print(image.size)
|
||||
|
||||
# tokens to text embeds
|
||||
prompt_embeds_list = []
|
||||
#for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
for tok, text_encoder in zip((tok1, tok2), self.text_encoders):
|
||||
prompt_embeds_out = text_encoder(
|
||||
tok.to(text_encoder.device),
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
print("prompt_embeds_out:")
|
||||
print(prompt_embeds_out.shape)
|
||||
|
||||
pooled_prompt_embeds = prompt_embeds_out[0]
|
||||
prompt_embeds = prompt_embeds_out.hidden_states[-2]
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
|
||||
prompt_embeds_list.append(prompt_embeds)
|
||||
|
||||
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
|
||||
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
|
||||
|
||||
print("prompt_embeds:")
|
||||
print(prompt_embeds.shape)
|
||||
print("pooled_prompt_embeds:")
|
||||
print(pooled_prompt_embeds.shape)
|
||||
|
||||
def save_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
|
||||
assert (
|
||||
self.train_ids is not None
|
||||
), "Initialize new tokens before saving embeddings."
|
||||
tensors = {}
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
|
||||
0
|
||||
] == len(self.tokenizers[0]), "Tokenizers should be the same."
|
||||
new_token_embeddings = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids
|
||||
]
|
||||
)
|
||||
tensors[txt_encoder_keys[idx]] = new_token_embeddings
|
||||
|
||||
save_file(tensors, file_path)
|
||||
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return self.text_encoders[0].dtype
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return self.text_encoders[0].device
|
||||
|
||||
def _compute_off_ratio(self, idx):
|
||||
# compute the off-std-ratio for the embeddings
|
||||
|
||||
text_encoder = self.text_encoders[idx]
|
||||
tokenizer = self.tokenizers[idx]
|
||||
|
||||
if text_encoder is None:
|
||||
off_ratio = -1
|
||||
else:
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
|
||||
index_updates = ~index_no_updates
|
||||
new_embeddings = (text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates])
|
||||
|
||||
off_ratio = std_token_embedding / new_embeddings.std()
|
||||
|
||||
return off_ratio
|
||||
|
||||
def fix_embedding_std(self, off_ratio_power = 0.1):
|
||||
std_penalty = 0.0
|
||||
idx = 0
|
||||
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if text_encoder is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
|
||||
index_updates = ~index_no_updates
|
||||
|
||||
new_embeddings = (text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates])
|
||||
|
||||
off_ratio = self._compute_off_ratio(idx)
|
||||
std_penalty += (off_ratio - 1.0)**2
|
||||
|
||||
if (off_ratio < 0.95) or (off_ratio > 1.05):
|
||||
print(f"std-off ratio-{idx} (target-std / embedding-std) = {off_ratio:.4f}, prob not ideal...")
|
||||
print(f"std_token_embedding: {std_token_embedding}")
|
||||
print(f"std new_embeddings: {new_embeddings.std()}")
|
||||
|
||||
# rescale the embeddings to have a more similar std as before:
|
||||
new_embeddings = new_embeddings * (off_ratio**off_ratio_power)
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
index_updates
|
||||
] = new_embeddings
|
||||
|
||||
idx += 1
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def retract_embeddings(self, print_stds = False):
|
||||
idx = 0
|
||||
means, stds = [], []
|
||||
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if text_encoder is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
index_no_updates
|
||||
] = (
|
||||
self.embeddings_settings[f"original_embeddings_{idx}"][index_no_updates]
|
||||
.to(device=text_encoder.device)
|
||||
.to(dtype=text_encoder.dtype)
|
||||
)
|
||||
|
||||
# for the parts that were updated, we can normalize them a bit
|
||||
# to have the same std as before
|
||||
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
|
||||
|
||||
index_updates = ~index_no_updates
|
||||
new_embeddings = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
index_updates
|
||||
]
|
||||
)
|
||||
|
||||
idx += 1
|
||||
|
||||
if 0:
|
||||
# get the actual embeddings that will get updated:
|
||||
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
inu[self.train_ids] = False
|
||||
updateable_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[~inu].detach().clone().to(dtype=torch.float32).cpu().numpy()
|
||||
|
||||
mean_0, mean_1 = updateable_embeddings[0].mean(), updateable_embeddings[1].mean()
|
||||
std_0, std_1 = updateable_embeddings[0].std(), updateable_embeddings[1].std()
|
||||
|
||||
means.append((mean_0, mean_1))
|
||||
stds.append((std_0, std_1))
|
||||
|
||||
if print_stds:
|
||||
print(f"Text Encoder {idx} token embeddings:")
|
||||
print(f" --- Means: ({mean_0:.6f}, {mean_1:.6f})")
|
||||
print(f" --- Stds: ({std_0:.6f}, {std_1:.6f})")
|
||||
|
||||
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
|
||||
# Assuming new tokens are of the format <s_i>
|
||||
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
|
||||
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
|
||||
tokenizer.add_special_tokens(special_tokens_dict)
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
|
||||
assert self.train_ids is not None, "New tokens could not be converted to IDs."
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids
|
||||
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
|
||||
|
||||
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
|
||||
if not os.path.exists(file_path):
|
||||
file_path = file_path.replace(".pti", ".safetensors")
|
||||
if not os.path.exists(file_path):
|
||||
raise FileNotFoundError(f"{file_path} does not exist.")
|
||||
|
||||
with safe_open(file_path, framework="pt", device=self.device.type) as f:
|
||||
for idx in range(len(self.text_encoders)):
|
||||
text_encoder = self.text_encoders[idx]
|
||||
tokenizer = self.tokenizers[idx]
|
||||
if text_encoder is None:
|
||||
continue
|
||||
try:
|
||||
loaded_embeddings = f.get_tensor(txt_encoder_keys[idx])
|
||||
except:
|
||||
loaded_embeddings = f.get_tensor(f"text_encoders_{idx}")
|
||||
self._load_embeddings(loaded_embeddings, tokenizer, text_encoder)
|
||||
@@ -1,89 +0,0 @@
|
||||
import os, json
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
from typing import Dict
|
||||
from peft import PeftModel
|
||||
from dataset_and_utils import TokenEmbeddingsHandler
|
||||
from safetensors.torch import save_file
|
||||
|
||||
'''
|
||||
from diffusers.utils import (
|
||||
convert_all_state_dict_to_peft,
|
||||
convert_state_dict_to_diffusers,
|
||||
convert_unet_state_dict_to_peft
|
||||
)
|
||||
'''
|
||||
|
||||
def patch_pipe_with_lora(pipe, lora_path):
|
||||
"""
|
||||
update the pipe with the lora model and the token embeddings
|
||||
"""
|
||||
|
||||
pipe.unet = PeftModel.from_pretrained(pipe.unet, lora_path)
|
||||
pipe.unet.merge_adapter()
|
||||
|
||||
# Load the textual_inversion token embeddings into the pipeline:
|
||||
try: #SDXL
|
||||
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
|
||||
except: #SD15
|
||||
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
|
||||
|
||||
embeddings_path = [f for f in os.listdir(lora_path) if f.endswith("embeddings.safetensors")][0]
|
||||
handler.load_embeddings(os.path.join(lora_path, embeddings_path))
|
||||
|
||||
return pipe
|
||||
|
||||
|
||||
def unet_attn_processors_state_dict(unet) -> Dict[str, torch.tensor]:
|
||||
"""
|
||||
Returns:
|
||||
a state dict containing just the attention processor parameters.
|
||||
"""
|
||||
attn_processors = unet.attn_processors
|
||||
|
||||
attn_processors_state_dict = {}
|
||||
|
||||
for attn_processor_key, attn_processor in attn_processors.items():
|
||||
for parameter_key, parameter in attn_processor.state_dict().items():
|
||||
attn_processors_state_dict[
|
||||
f"{attn_processor_key}.{parameter_key}"
|
||||
] = parameter
|
||||
|
||||
return attn_processors_state_dict
|
||||
|
||||
|
||||
def save_lora(output_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names):
|
||||
"""
|
||||
Save the LORA model to output_dir, optionally with some example images
|
||||
|
||||
"""
|
||||
print(f"Saving checkpoint at step.. {global_step}")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
args_dict["n_training_steps"] = global_step
|
||||
args_dict["total_n_imgs_seen"] = global_step * args_dict["train_batch_size"]
|
||||
|
||||
if not is_lora:
|
||||
lora_tensors = {
|
||||
name: param
|
||||
for name, param in unet.named_parameters()
|
||||
if name in unet_param_to_optimize_names
|
||||
}
|
||||
save_file(lora_tensors, f"{output_dir}/unet.safetensors",)
|
||||
elif len(unet_lora_parameters) > 0:
|
||||
unet.save_pretrained(save_directory = output_dir)
|
||||
|
||||
try:
|
||||
concept_name = args_dict["name"].lower()
|
||||
except:
|
||||
concept_name = "eden_concept_lora"
|
||||
|
||||
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
|
||||
concept_name = concept_name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
|
||||
|
||||
embedding_handler.save_embeddings(f"{output_dir}/{concept_name}_embeddings.safetensors",)
|
||||
|
||||
with open(f"{output_dir}/special_params.json", "w") as f:
|
||||
json.dump(token_dict, f)
|
||||
with open(f"{output_dir}/training_args.json", "w") as f:
|
||||
json.dump(args_dict, f, indent=4)
|
||||
@@ -0,0 +1,656 @@
|
||||
import fnmatch
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
import shutil
|
||||
import gc
|
||||
import numpy as np
|
||||
import argparse
|
||||
import itertools
|
||||
import zipfile
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
from tqdm import tqdm
|
||||
|
||||
from typing import Union, Iterable, List, Dict, Tuple, Optional, cast
|
||||
#from diffusers.training_utils import cast_training_params
|
||||
|
||||
from trainer.utils.utils import *
|
||||
from trainer.checkpoint import save_checkpoint
|
||||
from trainer.embedding_handler import TokenEmbeddingsHandler
|
||||
from trainer.dataset import PreprocessedDataset
|
||||
from trainer.config import TrainingConfig
|
||||
from trainer.models import print_trainable_parameters, load_models
|
||||
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer
|
||||
from trainer.inference import render_images, get_conditioning_signals
|
||||
from trainer.preprocess import preprocess
|
||||
from trainer.utils.io import make_validation_img_grid
|
||||
|
||||
from trainer.optimizer import (
|
||||
OptimizerCollection,
|
||||
get_optimizer_and_peft_models_text_encoder_lora,
|
||||
get_textual_inversion_optimizer,
|
||||
get_unet_lora_parameters,
|
||||
get_unet_optimizer
|
||||
)
|
||||
|
||||
def train(config: TrainingConfig):
|
||||
|
||||
seed_everything(config.seed)
|
||||
weight_dtype = dtype_map[config.weight_type]
|
||||
(
|
||||
pipe,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
noise_scheduler,
|
||||
text_encoder_one,
|
||||
text_encoder_two,
|
||||
vae,
|
||||
unet,
|
||||
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
|
||||
|
||||
from trainer.ti_cross_attn_loss import init_daam_loss
|
||||
|
||||
pipe, daam_loss = init_daam_loss(
|
||||
pipeline=pipe
|
||||
)
|
||||
config.sd_model_version = sd_model_version
|
||||
config.pretrained_model["version"] = sd_model_version
|
||||
|
||||
config, input_dir = preprocess(
|
||||
config,
|
||||
working_directory=config.output_dir,
|
||||
concept_mode=config.concept_mode,
|
||||
input_zip_path=config.lora_training_urls,
|
||||
caption_text=config.caption_prefix,
|
||||
mask_target_prompts=config.mask_target_prompts,
|
||||
target_size=config.resolution,
|
||||
crop_based_on_salience=config.crop_based_on_salience,
|
||||
use_face_detection_instead=config.use_face_detection_instead,
|
||||
left_right_flip_augmentation=config.left_right_flip_augmentation,
|
||||
augment_imgs_up_to_n = config.augment_imgs_up_to_n,
|
||||
caption_model = config.caption_model,
|
||||
seed = config.seed,
|
||||
)
|
||||
|
||||
if config.allow_tf32:
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
# Initialize new tokens for training.
|
||||
embedding_handler = TokenEmbeddingsHandler(
|
||||
text_encoders = [text_encoder_one, text_encoder_two],
|
||||
tokenizers = [tokenizer_one, tokenizer_two]
|
||||
)
|
||||
|
||||
embedding_handler.initialize_new_tokens(
|
||||
inserting_toks=config.inserting_list_tokens,
|
||||
starting_toks=None,
|
||||
seed=config.seed
|
||||
)
|
||||
|
||||
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
|
||||
embedding_handler.make_embeddings_trainable()
|
||||
embedding_handler.token_regularizer = ConditioningRegularizer(config, embedding_handler)
|
||||
embedding_handler.pre_optimize_token_embeddings(config, pipe)
|
||||
|
||||
# Turn off all gradients for now:
|
||||
unet.requires_grad_(False)
|
||||
vae.requires_grad_(False)
|
||||
text_encoders = embedding_handler.text_encoders
|
||||
for txt_encoder in text_encoders:
|
||||
if txt_encoder is not None:
|
||||
txt_encoder.requires_grad_(False)
|
||||
|
||||
if config.text_encoder_lora_optimizer is not None:
|
||||
print("Creating LoRA for text encoder...")
|
||||
optimizer_text_encoder_lora , text_encoder_peft_models = get_optimizer_and_peft_models_text_encoder_lora(
|
||||
text_encoders=text_encoders,
|
||||
lora_rank = config.text_encoder_lora_rank,
|
||||
lora_alpha_multiplier = config.lora_alpha_multiplier,
|
||||
use_dora = config.use_dora,
|
||||
optimizer_name = config.text_encoder_lora_optimizer,
|
||||
lora_lr = config.text_encoder_lora_lr,
|
||||
weight_decay = config.text_encoder_lora_weight_decay
|
||||
)
|
||||
else:
|
||||
optimizer_text_encoder_lora = None
|
||||
text_encoder_peft_models = [None] * len(text_encoders)
|
||||
|
||||
|
||||
embedding_handler.make_embeddings_trainable()
|
||||
if not config.disable_ti:
|
||||
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
|
||||
text_encoders=text_encoders,
|
||||
textual_inversion_lr=config.ti_lr,
|
||||
textual_inversion_weight_decay=config.ti_weight_decay,
|
||||
optimizer_name=config.ti_optimizer ## hardcoded
|
||||
)
|
||||
else:
|
||||
optimizer_ti = None
|
||||
textual_inversion_params = None
|
||||
|
||||
if not config.is_lora: # This code pathway has not been tested in a long while
|
||||
print(f"Doing full fine-tuning on the U-Net")
|
||||
unet.requires_grad_(True)
|
||||
unet_lora_parameters = None
|
||||
optimizer_text_encoder_lora = None
|
||||
unet_trainable_params = unet.parameters()
|
||||
else:
|
||||
# Do lora-training instead.
|
||||
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
|
||||
# target_blocks=["block"] for original IP-Adapter
|
||||
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
|
||||
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
|
||||
|
||||
unet, unet_trainable_params, unet_lora_parameters = get_unet_lora_parameters(
|
||||
lora_rank = config.lora_rank,
|
||||
lora_alpha_multiplier = config.lora_alpha_multiplier,
|
||||
lora_weight_decay=config.lora_weight_decay,
|
||||
use_dora = config.use_dora,
|
||||
unet=unet,
|
||||
pipe=pipe
|
||||
)
|
||||
|
||||
optimizer_unet = get_unet_optimizer(
|
||||
prodigy_d_coef=config.prodigy_d_coef,
|
||||
prodigy_growth_factor=config.unet_prodigy_growth_factor,
|
||||
lora_weight_decay=config.lora_weight_decay,
|
||||
use_dora=config.use_dora,
|
||||
unet_trainable_params=unet_trainable_params,
|
||||
optimizer_name=config.unet_optimizer_type
|
||||
)
|
||||
|
||||
print_trainable_parameters(unet, model_name = 'unet')
|
||||
for i, text_encoder in enumerate(text_encoders):
|
||||
if text_encoder is not None:
|
||||
print_trainable_parameters(text_encoder, model_name = f'text_encoder_{i}')
|
||||
|
||||
train_dataset = PreprocessedDataset(
|
||||
input_dir,
|
||||
pipe,
|
||||
vae.float(),
|
||||
size = config.train_img_size,
|
||||
do_cache=config.do_cache,
|
||||
substitute_caption_map=config.token_dict,
|
||||
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
|
||||
train_batch_size=config.train_batch_size
|
||||
)
|
||||
# offload the vae to cpu:
|
||||
vae = vae.to('cpu')
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset,
|
||||
batch_size=config.train_batch_size,
|
||||
shuffle=True,
|
||||
num_workers=config.dataloader_num_workers,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps)
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
|
||||
|
||||
if config.max_train_steps is None:
|
||||
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
|
||||
|
||||
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
|
||||
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
|
||||
|
||||
print(f"--- Num samples = {len(train_dataset)}")
|
||||
print(f"--- Num batches each epoch = {len(train_dataloader)}")
|
||||
print(f"--- Num Epochs = {config.num_train_epochs}")
|
||||
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
|
||||
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
|
||||
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
|
||||
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
|
||||
|
||||
global_step = 0
|
||||
last_save_step = 0
|
||||
|
||||
progress_bar = tqdm(range(global_step, config.max_train_steps), position=0, leave=True)
|
||||
checkpoint_dir = os.path.join(str(config.output_dir), "checkpoints")
|
||||
if os.path.exists(checkpoint_dir):
|
||||
shutil.rmtree(checkpoint_dir)
|
||||
os.makedirs(f"{checkpoint_dir}")
|
||||
|
||||
# Data tracking inits:
|
||||
start_time, images_done = time.time(), 0
|
||||
prompt_embeds_norms = {'main':[], 'reg':[]}
|
||||
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': []}
|
||||
grad_norms, token_stds = {'unet': []}, {}
|
||||
for i in range(len(text_encoders)):
|
||||
grad_norms[f'text_encoder_{i}'] = []
|
||||
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
|
||||
|
||||
# default value of cold (pre-warmup) optimizer lr:
|
||||
if config.sd_model_version == "sdxl":
|
||||
if config.is_lora: # let textual_inversion do the work first!
|
||||
base_lr = 1.0e-5
|
||||
else:
|
||||
base_lr = 3.0e-5
|
||||
elif config.sd_model_version == "sd15":
|
||||
# let lora training kick in soonish (pure ti for sd15 is not working super well in my tests)
|
||||
base_lr = 1.0e-4
|
||||
|
||||
#######################################################################################################
|
||||
|
||||
"""
|
||||
Storing all optimizers in a single container
|
||||
"""
|
||||
optimizer_collection = OptimizerCollection(
|
||||
optimizer_textual_inversion=optimizer_ti,
|
||||
optimizer_text_encoders=optimizer_text_encoder_lora,
|
||||
optimizer_unet=optimizer_unet,
|
||||
debug = config.debug
|
||||
)
|
||||
optimizers = optimizer_collection.optimizers
|
||||
|
||||
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
|
||||
|
||||
for epoch in range(config.num_train_epochs):
|
||||
if config.aspect_ratio_bucketing:
|
||||
train_dataset.bucket_manager.start_epoch()
|
||||
progress_bar.set_description(f"# Trainer step: {global_step}, epoch: {epoch}")
|
||||
|
||||
for step, batch in enumerate(train_dataloader):
|
||||
progress_bar.update(1)
|
||||
finegrained_epoch = epoch + step / len(train_dataloader)
|
||||
completion_f = finegrained_epoch / config.num_train_epochs
|
||||
|
||||
# param_groups[1] goes from ti_lr to 0.0 over the course of training
|
||||
if config.ti_optimizer != "prodigy": # Update ti_learning rate gradually:
|
||||
if optimizers['textual_inversion'] is not None:
|
||||
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0
|
||||
# warmup the ti-lr:
|
||||
if config.ti_lr_warmup_steps > 0:
|
||||
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0)
|
||||
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
|
||||
if config.freeze_ti_after_completion_f <= completion_f:
|
||||
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
|
||||
|
||||
if optimizers['text_encoders'] is not None:
|
||||
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
|
||||
|
||||
# warmup the txt-encoder lr:
|
||||
if config.txt_encoders_lr_warmup_steps > 0 and optimizers['text_encoders'] is not None:
|
||||
warmup_f = min(global_step / config.txt_encoders_lr_warmup_steps, 1.0)
|
||||
optimizers['text_encoders'].param_groups[0]['lr'] *= warmup_f
|
||||
|
||||
if optimizers['unet'] is not None:
|
||||
# Calculate the exponential factor
|
||||
exp_factor = (config.unet_lr / base_lr) ** (global_step / config.unet_lr_warmup_steps)
|
||||
# Apply the exponential learning rate
|
||||
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
|
||||
|
||||
if not config.aspect_ratio_bucketing:
|
||||
captions, vae_latent, mask = batch
|
||||
else:
|
||||
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
|
||||
|
||||
captions = list(captions)
|
||||
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
|
||||
config, pipe, captions
|
||||
)
|
||||
|
||||
# Sample noise that we'll add to the latents:
|
||||
vae_latent = vae_latent.to(weight_dtype)
|
||||
noise = torch.randn_like(vae_latent)
|
||||
|
||||
if config.noise_offset > 0.0:
|
||||
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
|
||||
noise += config.noise_offset * torch.randn(
|
||||
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
|
||||
|
||||
timesteps = torch.randint(
|
||||
0,
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
(vae_latent.shape[0],),
|
||||
device=vae_latent.device,
|
||||
).long()
|
||||
|
||||
noisy_latent = noise_scheduler.add_noise(vae_latent, noise, timesteps)
|
||||
|
||||
# Predict the noise residual
|
||||
model_pred = unet(
|
||||
noisy_latent,
|
||||
timesteps,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep_cond=None,
|
||||
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
"""
|
||||
distirbution shift loss
|
||||
"""
|
||||
non_ti_heatmaps = []
|
||||
ti_heatmaps = []
|
||||
ti_token_indices = [0,1]
|
||||
batch_index = 0
|
||||
|
||||
token_strings = [
|
||||
pipe.tokenizer.decode(x)
|
||||
for x in pipe.tokenizer.encode(captions[batch_index])
|
||||
]
|
||||
|
||||
for text_token_index in range(1, len(token_strings)-1):
|
||||
if text_token_index in ti_token_indices:
|
||||
ti_heatmaps.append(
|
||||
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
|
||||
)
|
||||
else:
|
||||
# we unsqueeze because we'll stack them together and then calculate the min, max and the mean
|
||||
non_ti_heatmaps.append(
|
||||
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
|
||||
)
|
||||
|
||||
non_ti_heatmaps = torch.cat(
|
||||
non_ti_heatmaps,
|
||||
dim = 0
|
||||
)
|
||||
ti_heatmaps = torch.cat(
|
||||
ti_heatmaps,
|
||||
dim = 0
|
||||
)
|
||||
|
||||
non_ti_dist = {
|
||||
"mean": non_ti_heatmaps.mean(),
|
||||
"min": non_ti_heatmaps.min(),
|
||||
"max": non_ti_heatmaps.min()
|
||||
}
|
||||
|
||||
ti_dist = {
|
||||
"mean": ti_heatmaps.mean(),
|
||||
"min": ti_heatmaps.min(),
|
||||
"max": ti_heatmaps.min()
|
||||
}
|
||||
|
||||
dist_loss = (non_ti_dist["mean"] - ti_dist["mean"].to(non_ti_dist["mean"].device)) ** 2
|
||||
|
||||
|
||||
if global_step % 20 == 0:
|
||||
batch_index = 0
|
||||
folder = "./heatmaps"
|
||||
fig = plt.figure()
|
||||
|
||||
token_strings = [
|
||||
pipe.tokenizer.decode(x)
|
||||
for x in pipe.tokenizer.encode(captions[batch_index])
|
||||
]
|
||||
|
||||
|
||||
plot_token_indices = range(len(token_strings))
|
||||
|
||||
fig, ax = plt.subplots(nrows=1, ncols=len(plot_token_indices), figsize = (int(3 * len(plot_token_indices)) , 10))
|
||||
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
|
||||
im = ax[idx].imshow(heatmap)
|
||||
ax[idx].set_title(f"{token_strings[text_token_index]}\n timestep: {timesteps[batch_index].item()}\nmax: {heatmap.max().item()}\nmin: {heatmap.min().item()}\nnorm: {heatmap.norm().item()}")
|
||||
ax[idx].axis("off")
|
||||
|
||||
fig.savefig(
|
||||
os.path.join(
|
||||
folder,
|
||||
f"{global_step}.jpg"
|
||||
)
|
||||
)
|
||||
plt.close(fig)
|
||||
|
||||
"""
|
||||
histogram to visualize the distributions of the cross attention values for each text token on the image space
|
||||
"""
|
||||
fig = plt.figure()
|
||||
fig.suptitle(f"Dist loss: {dist_loss.item()}")
|
||||
plot_token_indices = range(1, len(token_strings)-1)
|
||||
for idx, text_token_index in enumerate(plot_token_indices):
|
||||
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
|
||||
plt.hist(heatmap.reshape(-1), bins = 30, label = token_strings[text_token_index], alpha = 0.5)
|
||||
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
|
||||
|
||||
plt.xlabel("Value")
|
||||
plt.ylabel("Number of instances")
|
||||
plt.grid()
|
||||
# Adjust the layout to prevent the legend from being cut off
|
||||
plt.tight_layout()
|
||||
fig.savefig(
|
||||
os.path.join(
|
||||
folder,
|
||||
f"{global_step}_heatmap.jpg"
|
||||
),
|
||||
bbox_inches='tight' # This ensures the legend is not cut off when saving
|
||||
)
|
||||
plt.close(fig) # Close the figure to free up memory
|
||||
|
||||
# Compute the loss:
|
||||
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
|
||||
losses['img_loss'].append(loss.item())
|
||||
|
||||
if config.training_attributes["gpt_description"] and config.debug:
|
||||
concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe)
|
||||
# Dont apply this loss, just plot it for now:
|
||||
loss += 0.0 * concept_description_loss
|
||||
losses['concept_description_loss'].append(concept_description_loss.item())
|
||||
|
||||
if config.l1_penalty > 0.0 and unet_lora_parameters:
|
||||
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
|
||||
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
|
||||
loss += config.l1_penalty * l1_norm
|
||||
|
||||
if optimizers['textual_inversion'] is not None and optimizers['textual_inversion'].param_groups[0]['lr'] > 0.0:
|
||||
loss, losses, prompt_embeds_norms = embedding_handler.token_regularizer.apply_regularization(loss, losses, prompt_embeds_norms, prompt_embeds, pipe = pipe)
|
||||
|
||||
losses['tot_loss'].append(loss.item())
|
||||
loss = loss + 1e-4 * dist_loss
|
||||
loss = loss / config.gradient_accumulation_steps
|
||||
loss.backward()
|
||||
|
||||
last_batch = (step + 1 == len(train_dataloader))
|
||||
if (step + 1) % config.gradient_accumulation_steps == 0 or last_batch:
|
||||
|
||||
if optimizers['textual_inversion'] is not None:
|
||||
# zero out the gradients of the non-trained text-encoder embeddings
|
||||
for i, embedding_tensor in enumerate(textual_inversion_params):
|
||||
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
|
||||
|
||||
if config.debug:
|
||||
# Track the average gradient norms:
|
||||
grad_norms['unet'].append(compute_grad_norm(itertools.chain(unet.parameters())).item())
|
||||
for i, text_encoder in enumerate(text_encoders):
|
||||
if text_encoder is not None:
|
||||
text_encoder_norm = compute_grad_norm(itertools.chain(text_encoder.parameters())).item()
|
||||
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
|
||||
|
||||
optimizer_collection.step()
|
||||
optimizer_collection.zero_grad()
|
||||
|
||||
#############################################################################################################
|
||||
|
||||
if config.debug:
|
||||
# Track the token embedding stds:
|
||||
trainable_embeddings, _ = embedding_handler.get_trainable_embeddings()
|
||||
for idx in range(len(text_encoders)):
|
||||
if text_encoders[idx] is not None:
|
||||
embedding_stds = trainable_embeddings[f'txt_encoder_{idx}'].detach().float().std(dim=1)
|
||||
for std_i, std in enumerate(embedding_stds):
|
||||
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
|
||||
|
||||
# Print some statistics:
|
||||
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0:
|
||||
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
||||
os.makedirs(output_save_dir, exist_ok=True)
|
||||
config.save_as_json(
|
||||
os.path.join(output_save_dir, "training_args.json")
|
||||
)
|
||||
save_checkpoint(
|
||||
output_dir=output_save_dir,
|
||||
global_step=global_step,
|
||||
unet=unet,
|
||||
embedding_handler=embedding_handler,
|
||||
token_dict=config.token_dict,
|
||||
is_lora=config.is_lora,
|
||||
unet_lora_parameters=unet_lora_parameters,
|
||||
name=config.name,
|
||||
text_encoder_peft_models=text_encoder_peft_models,
|
||||
pretrained_model_version=config.pretrained_model["version"]
|
||||
)
|
||||
last_save_step = global_step
|
||||
|
||||
if config.debug:
|
||||
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings()
|
||||
for idx, text_encoder in enumerate(text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
n = len(token_embeddings[f'txt_encoder_{idx}'])
|
||||
for i in range(n):
|
||||
token = trainable_tokens[f'txt_encoder_{idx}'][i]
|
||||
# Strip any backslashes from the token name:
|
||||
token = token.replace("/", "_")
|
||||
embedding = token_embeddings[f'txt_encoder_{idx}'][i]
|
||||
plot_torch_hist(embedding, global_step, os.path.join(config.output_dir, 'ti_embeddings') , f"enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'red')
|
||||
|
||||
embedding_handler.print_token_info()
|
||||
if config.is_lora: # plotting this hist for full unet parameters can run OOM
|
||||
plot_torch_hist(unet_lora_parameters, global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
|
||||
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
|
||||
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
|
||||
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
|
||||
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
|
||||
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
|
||||
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
|
||||
|
||||
validation_prompts = render_images(
|
||||
pipe = pipe,
|
||||
render_size = config.validation_img_size,
|
||||
lora_path = output_save_dir,
|
||||
train_step = global_step,
|
||||
seed = config.seed,
|
||||
is_lora = config.is_lora,
|
||||
pretrained_model = config.pretrained_model,
|
||||
lora_scale = config.sample_imgs_lora_scale,
|
||||
n_imgs = config.n_sample_imgs,
|
||||
device = config.device,
|
||||
checkpoint_folder = None
|
||||
)
|
||||
img_grid_path = make_validation_img_grid(output_save_dir)
|
||||
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
images_done += config.train_batch_size
|
||||
global_step += 1
|
||||
|
||||
if global_step % (config.max_train_steps//50) == 0:
|
||||
progress = (global_step / config.max_train_steps) + 0.05
|
||||
#print_system_info()
|
||||
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
|
||||
yield np.min((progress, 1.0))
|
||||
|
||||
if global_step > config.max_train_steps:
|
||||
print("Reached max steps, stopping training!", flush = True)
|
||||
break
|
||||
|
||||
# final_save
|
||||
if (global_step - last_save_step) > 26:
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
||||
else:
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
|
||||
|
||||
if config.debug:
|
||||
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
|
||||
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
|
||||
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
|
||||
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
|
||||
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
|
||||
|
||||
if not os.path.exists(output_save_dir):
|
||||
os.makedirs(output_save_dir, exist_ok=True)
|
||||
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
|
||||
save_checkpoint(
|
||||
output_dir=output_save_dir,
|
||||
global_step=global_step,
|
||||
unet=unet,
|
||||
embedding_handler=embedding_handler,
|
||||
token_dict=config.token_dict,
|
||||
is_lora=config.is_lora,
|
||||
unet_lora_parameters=unet_lora_parameters,
|
||||
name=config.name,
|
||||
pretrained_model_version=config.pretrained_model["version"]
|
||||
)
|
||||
|
||||
if config.debug and 0:
|
||||
# Reload the entire pipe from disk + LoRa:
|
||||
pipe_to_use = None
|
||||
checkpoint_folder = output_save_dir
|
||||
del unet
|
||||
del vae
|
||||
del text_encoder_one
|
||||
del text_encoder_two
|
||||
del tokenizer_one
|
||||
del tokenizer_two
|
||||
del embedding_handler
|
||||
del pipe
|
||||
del train_dataloader
|
||||
del train_dataset
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
else:
|
||||
# Just render images with the active pipe (faster, easier):
|
||||
pipe_to_use = pipe
|
||||
checkpoint_folder = None
|
||||
|
||||
validation_prompts = render_images(
|
||||
pipe = pipe_to_use,
|
||||
render_size=config.validation_img_size,
|
||||
lora_path=output_save_dir,
|
||||
train_step=global_step,
|
||||
seed=config.seed,
|
||||
is_lora=config.is_lora,
|
||||
pretrained_model=config.pretrained_model,
|
||||
lora_scale=config.sample_imgs_lora_scale,
|
||||
n_imgs = config.n_sample_imgs,
|
||||
n_steps = 30,
|
||||
device = config.device,
|
||||
checkpoint_folder=checkpoint_folder
|
||||
)
|
||||
|
||||
img_grid_path = make_validation_img_grid(output_save_dir)
|
||||
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
|
||||
|
||||
else:
|
||||
print(f"Skipping final save, {output_save_dir} already exists")
|
||||
|
||||
if config.debug:
|
||||
# Create a zipfile of all the *.py files in the directory
|
||||
parent_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
zip_file_path = os.path.join(config.output_dir, 'source_code.zip')
|
||||
with zipfile.ZipFile(zip_file_path, 'w', zipfile.ZIP_DEFLATED) as zipf:
|
||||
zipdir(parent_dir, zipf)
|
||||
|
||||
config.job_time = time.time() - config.start_time
|
||||
config.training_attributes["validation_prompts"] = validation_prompts
|
||||
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
|
||||
print("Training job complete, saving outputs...", flush = True)
|
||||
print("------------------------------------------")
|
||||
|
||||
return config, output_save_dir
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Train a concept')
|
||||
parser.add_argument('config_filename', type=str, help='Input JSON configuration file')
|
||||
args = parser.parse_args()
|
||||
|
||||
config = TrainingConfig.from_json(file_path=args.config_filename)
|
||||
|
||||
print("Starting new LoRa training run with config:")
|
||||
print(config)
|
||||
print("------------------------------------------")
|
||||
|
||||
for progress in train(config=config):
|
||||
print(f"Progress: {(100*progress):.2f}%", end="\r")
|
||||
|
||||
print("Training done :)")
|
||||
@@ -0,0 +1,130 @@
|
||||
import os
|
||||
import tarfile
|
||||
import json
|
||||
import time
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from main import train
|
||||
from trainer.config import TrainingConfig, model_paths
|
||||
from trainer.utils.io import clean_filename
|
||||
|
||||
import folder_paths
|
||||
import comfy.utils
|
||||
|
||||
class Eden_LoRa_trainer:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"training_images_folder_path": ("STRING", {"default": "."}),
|
||||
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||
"lora_name": ("STRING", {"default": "Eden_LoRa"}),
|
||||
"mode": (["style", "face", "object"], ),
|
||||
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
|
||||
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
|
||||
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
|
||||
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
|
||||
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
|
||||
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
|
||||
"use_dora": ("BOOLEAN", {"default": False}),
|
||||
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}),
|
||||
"debug_mode": ("BOOLEAN", {"default": False}),
|
||||
"checkpointing_steps": ("INT", {"default": 200, "min": 10, "max": 2000}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "Eden 🌱"
|
||||
RETURN_TYPES = ("IMAGE", "STRING", "STRING", "STRING")
|
||||
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
|
||||
FUNCTION = "train_lora"
|
||||
|
||||
def train_lora(self,
|
||||
training_images_folder_path,
|
||||
ckpt_name,
|
||||
lora_name = "eden_lora",
|
||||
mode = "style",
|
||||
seed = 0,
|
||||
resolution = 521,
|
||||
train_batch_size = 4,
|
||||
max_train_steps = 400,
|
||||
ti_lr = 0.001,
|
||||
unet_lr = 0.001,
|
||||
lora_rank = 16,
|
||||
use_dora = False,
|
||||
n_tokens = 2,
|
||||
debug_mode = False,
|
||||
checkpointing_steps = 1000,
|
||||
):
|
||||
|
||||
print("Starting new training job...")
|
||||
|
||||
# Overwrite hardcoded paths to point to comfyUI folders:
|
||||
model_paths.set_path("CLIP", os.path.join(folder_paths.models_dir, "clipseg"))
|
||||
model_paths.set_path("BLIP", os.path.join(folder_paths.models_dir, "blip"))
|
||||
model_paths.set_path("SR", os.path.join(folder_paths.models_dir, "upscale_models"))
|
||||
model_paths.set_path("SD", os.path.join(folder_paths.models_dir, "checkpoints"))
|
||||
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
|
||||
config = TrainingConfig(
|
||||
name=lora_name,
|
||||
lora_training_urls=training_images_folder_path,
|
||||
concept_mode=mode,
|
||||
ckpt_path=ckpt_path,
|
||||
seed=seed,
|
||||
resolution=resolution,
|
||||
train_batch_size=train_batch_size,
|
||||
max_train_steps=max_train_steps,
|
||||
checkpointing_steps=checkpointing_steps,
|
||||
ti_lr=ti_lr,
|
||||
unet_lr=unet_lr,
|
||||
lora_rank=lora_rank,
|
||||
use_dora=use_dora,
|
||||
caption_model="blip",
|
||||
n_tokens=n_tokens,
|
||||
verbose=True,
|
||||
debug=debug_mode,
|
||||
)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(100)
|
||||
|
||||
with torch.inference_mode(False):
|
||||
train_generator = train(config=config)
|
||||
while True:
|
||||
try:
|
||||
progress_f = next(train_generator)
|
||||
pbar.update_absolute(progress_f * 100)
|
||||
except StopIteration as e:
|
||||
config, output_save_dir = e.value # Capture the return value
|
||||
break
|
||||
|
||||
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
|
||||
|
||||
attributes = {}
|
||||
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
||||
attributes['job_time_seconds'] = config.job_time
|
||||
|
||||
print(f"LORA training node finished in {config.job_time:.1f} seconds")
|
||||
print("---------- Made with love by Eden.art 🌱 ----------")
|
||||
|
||||
# safetensors paths:
|
||||
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
|
||||
|
||||
# find the index of the path containing "_embeddings.safetensors":
|
||||
for i, path in enumerate(paths):
|
||||
if "_embeddings.safetensors" in path:
|
||||
embedding_path = path
|
||||
else:
|
||||
lora_path = path
|
||||
|
||||
# Load the grid image:
|
||||
grid_image = Image.open(validation_grid_img_path)
|
||||
grid_image = np.array(grid_image).astype(np.float32) / 255.0
|
||||
grid_image = torch.from_numpy(grid_image)[None,]
|
||||
|
||||
final_msg = f"LoRa trained in {config.job_time/60:.1f} minutes. Files saved at {output_save_dir}"
|
||||
|
||||
return (grid_image, lora_path, embedding_path, final_msg)
|
||||
+72
-353
@@ -7,16 +7,20 @@ import random
|
||||
import torch
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from collections import OrderedDict
|
||||
|
||||
from cog import BasePredictor, BaseModel, File, Input, Path as cogPath
|
||||
from dotenv import load_dotenv
|
||||
from preprocess import preprocess
|
||||
from trainer_pti import main
|
||||
from main import train
|
||||
from typing import Iterator, Optional
|
||||
from io_utils import MODEL_INFO, download_weights, clean_filename
|
||||
|
||||
from trainer.preprocess import preprocess
|
||||
from trainer.models import pretrained_models
|
||||
from trainer.config import TrainingConfig
|
||||
from trainer.utils.io import clean_filename
|
||||
from trainer.utils.utils import seed_everything
|
||||
|
||||
DEBUG_MODE = False
|
||||
XANDER_EXPERIMENT = False
|
||||
|
||||
load_dotenv()
|
||||
|
||||
@@ -25,6 +29,7 @@ os.environ["TRANSFORMERS_CACHE"] = "/src/.huggingface/"
|
||||
os.environ["DIFFUSERS_CACHE"] = "/src/.huggingface/"
|
||||
os.environ["HF_HOME"] = "/src/.huggingface/"
|
||||
|
||||
|
||||
class CogOutput(BaseModel):
|
||||
files: Optional[list[cogPath]] = []
|
||||
name: Optional[str] = None
|
||||
@@ -47,355 +52,93 @@ class Predictor(BasePredictor):
|
||||
default="unnamed"
|
||||
),
|
||||
lora_training_urls: str = Input(
|
||||
description="Training images for new LORA concept (can be image urls or a .zip file of images)",
|
||||
default=None
|
||||
description="Training images for new LORA concept (can be image urls or an url to a .zip file of images)"
|
||||
),
|
||||
concept_mode: str = Input(
|
||||
description=" 'face' / 'style' / 'object' (default)",
|
||||
default="object",
|
||||
description="What are you trying to learn?",
|
||||
choices=["style", "face", "object"],
|
||||
default="style",
|
||||
),
|
||||
sd_model_version: str = Input(
|
||||
description=" 'sdxl' / 'sd15' ",
|
||||
description="SDXL gives much better LoRa's if you just need static images. If you want to make AnimateDiff animations, train an SD15 lora.",
|
||||
choices=["sdxl", "sd15"],
|
||||
default="sdxl",
|
||||
),
|
||||
max_train_steps: int = Input(
|
||||
description="Number of training steps. Increasing this usually leads to overfitting, only viable if you have > 100 training imgs. For faces you may want to reduce to eg 300",
|
||||
default=400
|
||||
),
|
||||
resolution: int = Input(
|
||||
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
|
||||
default=512
|
||||
),
|
||||
train_batch_size: int = Input(
|
||||
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
|
||||
default=4
|
||||
),
|
||||
unet_lr: float = Input(
|
||||
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
|
||||
default=0.001
|
||||
),
|
||||
ti_lr: float = Input(
|
||||
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
|
||||
default=0.001
|
||||
),
|
||||
lora_rank: int = Input(
|
||||
description="Rank of LoRA embeddings for the unet.",
|
||||
default=16
|
||||
),
|
||||
use_dora: bool = Input(
|
||||
description="Use Dora instead of LoRa",
|
||||
default=False,
|
||||
),
|
||||
n_tokens: int = Input(
|
||||
description="How many new tokens to train (highly recommended to leave this at 2)",
|
||||
ge=1, le=3, default=2
|
||||
),
|
||||
seed: int = Input(
|
||||
description="Random seed for reproducible training. Leave empty to use a random seed",
|
||||
default=None,
|
||||
),
|
||||
resolution: int = Input(
|
||||
description="Square pixel resolution which your images will be resized to for training recommended [768-1024]",
|
||||
default=960,
|
||||
),
|
||||
train_batch_size: int = Input(
|
||||
description="Batch size (per device) for training",
|
||||
default=4,
|
||||
),
|
||||
num_train_epochs: int = Input(
|
||||
description="Number of epochs to loop through your training dataset",
|
||||
default=10000,
|
||||
),
|
||||
max_train_steps: int = Input(
|
||||
description="Number of individual training steps. Takes precedence over num_train_epochs",
|
||||
default=600,
|
||||
),
|
||||
checkpointing_steps: int = Input(
|
||||
description="Number of steps between saving checkpoints. Set to very very high number to disable checkpointing, because you don't need one.",
|
||||
default=10000,
|
||||
),
|
||||
gradient_accumulation_steps: int = Input(
|
||||
description="Number of training steps to accumulate before a backward pass. Effective batch size = gradient_accumulation_steps * batch_size",
|
||||
default=1,
|
||||
),
|
||||
is_lora: bool = Input(
|
||||
description="Whether to use LoRA training. If set to False, will use Full fine tuning",
|
||||
default=True,
|
||||
),
|
||||
prodigy_d_coef: float = Input(
|
||||
description="Multiplier for internal learning rate of Prodigy optimizer",
|
||||
default=0.5,
|
||||
),
|
||||
ti_lr: float = Input(
|
||||
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
|
||||
default=1e-3,
|
||||
),
|
||||
ti_weight_decay: float = Input(
|
||||
description="weight decay for textual inversion embeddings. Don't alter unless you know what you're doing.",
|
||||
default=3e-4,
|
||||
),
|
||||
lora_weight_decay: float = Input(
|
||||
description="weight decay for lora parameters. Don't alter unless you know what you're doing.",
|
||||
default=0.002,
|
||||
),
|
||||
l1_penalty: float = Input(
|
||||
description="Sparsity penalty for the LoRA matrices, possibly improves merge-ability and generalization",
|
||||
default=0.1,
|
||||
),
|
||||
lora_param_scaler: float = Input(
|
||||
description="Multiplier for the starting weights of the lora matrices",
|
||||
default=0.5,
|
||||
),
|
||||
snr_gamma: float = Input(
|
||||
description="see https://arxiv.org/pdf/2303.09556.pdf, set to None to disable snr training",
|
||||
default=5.0,
|
||||
),
|
||||
lora_rank: int = Input(
|
||||
description="Rank of LoRA embeddings. For faces 5 is good, for complex concepts / styles you can try 8 or 12",
|
||||
default=12,
|
||||
),
|
||||
caption_prefix: str = Input(
|
||||
description="Prefix text prepended to automatic captioning. Must contain the 'TOK'. Example is 'a photo of TOK, '. If empty, chatgpt will take care of this automatically",
|
||||
default="",
|
||||
),
|
||||
caption_model: str = Input(
|
||||
description="Which captioning model to use. ['gpt4-v', 'blip'] are supported right now",
|
||||
default="blip",
|
||||
),
|
||||
left_right_flip_augmentation: bool = Input(
|
||||
description="Add left-right flipped version of each img to the training data, recommended for most cases. If you are learning a face, you prob want to disable this",
|
||||
default=True,
|
||||
),
|
||||
augment_imgs_up_to_n: int = Input(
|
||||
description="Apply data augmentation (no lr-flipping) until there are n training samples (0 disables augmentation completely)",
|
||||
default=20,
|
||||
),
|
||||
n_tokens: int = Input(
|
||||
description="How many new tokens to inject per concept",
|
||||
default=2,
|
||||
),
|
||||
mask_target_prompts: str = Input(
|
||||
description="Prompt that describes most important part of the image, will be used for CLIP-segmentation. For example, if you are learning a person 'face' would be a good segmentation prompt",
|
||||
default=None,
|
||||
),
|
||||
crop_based_on_salience: bool = Input(
|
||||
description="If you want to crop the image to `target_size` based on the important parts of the image, set this to True. If you want to crop the image based on face detection, set this to False",
|
||||
default=True,
|
||||
),
|
||||
use_face_detection_instead: bool = Input(
|
||||
description="If you want to use face detection instead of CLIPSeg for masking. For face applications, we recommend using this option.",
|
||||
default=False,
|
||||
),
|
||||
clipseg_temperature: float = Input(
|
||||
description="How blurry you want the CLIPSeg mask to be. We recommend this value be something between `0.5` to `1.0`. If you want to have more sharp mask (but thus more errorful), you can decrease this value.",
|
||||
default=0.7,
|
||||
),
|
||||
verbose: bool = Input(description="verbose output", default=True),
|
||||
run_name: str = Input(
|
||||
description="Subdirectory where all files will be saved",
|
||||
default=str(int(time.time())),
|
||||
),
|
||||
debug: bool = Input(
|
||||
description="for debugging locally only (dont activate this on replicate)",
|
||||
default=False,
|
||||
),
|
||||
hard_pivot: bool = Input(
|
||||
description="Use hard freeze for ti_lr. If set to False, will use soft transition of learning rates",
|
||||
default=False,
|
||||
),
|
||||
off_ratio_power: float = Input(
|
||||
description="How strongly to correct the embedding std vs the avg-std (0=off, 0.05=weak, 0.1=standard)",
|
||||
default=0.1,
|
||||
),
|
||||
|
||||
) -> Iterator[GENERATOR_OUTPUT_TYPE]:
|
||||
|
||||
"""
|
||||
lambda @1024 training speed (SDXL):
|
||||
lambda training speed (SDXL):
|
||||
bs=2: 3.5 imgs/s, 1.8 batches/s
|
||||
bs=3: 5.1 imgs/s
|
||||
bs=4: 6.0 imgs/s,
|
||||
bs=6: 8.0 imgs/s,
|
||||
"""
|
||||
|
||||
start_time = time.time()
|
||||
out_root_dir = "lora_models"
|
||||
|
||||
if seed is None:
|
||||
seed = np.random.randint(0, 2**32 - 1)
|
||||
|
||||
# Try to make the training reproducible:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
if concept_mode == "face":
|
||||
left_right_flip_augmentation = False # always disable lr flips for face mode!
|
||||
mask_target_prompts = "face"
|
||||
clipseg_temperature = 0.4
|
||||
|
||||
if concept_mode == "concept": # gracefully catch any old versions of concept_mode
|
||||
concept_mode = "object"
|
||||
|
||||
if concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
|
||||
l1_penalty = 0.05
|
||||
|
||||
print(f"cog:predict:train_lora:{concept_mode}")
|
||||
debug = False
|
||||
|
||||
print("cog:predict starting new training job...")
|
||||
if not debug:
|
||||
yield CogOutput(name=name, progress=0.0)
|
||||
|
||||
# Initialize pretrained_model dictionary
|
||||
pretrained_model = {"version": sd_model_version}
|
||||
pretrained_model.update(MODEL_INFO[pretrained_model['version']])
|
||||
|
||||
# Download the weights if they don't exist locally
|
||||
if not os.path.exists(pretrained_model['path']):
|
||||
download_weights(pretrained_model['url'], pretrained_model['path'])
|
||||
|
||||
# hardcoded for now:
|
||||
token_list = [f"TOK:{n_tokens}"]
|
||||
#token_list = ["TOK1:2", "TOK2:2"]
|
||||
|
||||
token_dict = OrderedDict({})
|
||||
all_token_lists = []
|
||||
running_tok_cnt = 0
|
||||
for token in token_list:
|
||||
token_name, n_tok = token.split(":")
|
||||
n_tok = int(n_tok)
|
||||
special_tokens = [f"<s{i + running_tok_cnt}>" for i in range(n_tok)]
|
||||
token_dict[token_name] = "".join(special_tokens)
|
||||
all_token_lists.extend(special_tokens)
|
||||
running_tok_cnt += n_tok
|
||||
|
||||
if 0:
|
||||
# overwrite some settings for experimentation:
|
||||
lora_param_scaler = 0.1
|
||||
l1_penalty = 0.2
|
||||
prodigy_d_coef = 0.2
|
||||
ti_lr = 1e-3
|
||||
lora_rank = 24
|
||||
|
||||
lora_training_urls = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/plantoid_5.zip"
|
||||
concept_mode = "object"
|
||||
mask_target_prompts = ""
|
||||
left_right_flip_augmentation = True
|
||||
|
||||
output_dir1 = os.path.join(out_root_dir, run_name + "_xander")
|
||||
input_dir1, n_imgs1, trigger_text1, segmentation_prompt1, captions1 = preprocess(
|
||||
output_dir1,
|
||||
concept_mode,
|
||||
input_zip_path=lora_training_urls,
|
||||
caption_text=caption_prefix,
|
||||
mask_target_prompts=mask_target_prompts,
|
||||
target_size=resolution,
|
||||
crop_based_on_salience=crop_based_on_salience,
|
||||
use_face_detection_instead=use_face_detection_instead,
|
||||
temp=clipseg_temperature,
|
||||
left_right_flip_augmentation=left_right_flip_augmentation,
|
||||
augment_imgs_up_to_n = augment_imgs_up_to_n,
|
||||
seed = seed,
|
||||
caption_model = caption_model
|
||||
)
|
||||
|
||||
lora_training_urls = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene_5.zip"
|
||||
concept_mode = "face"
|
||||
mask_target_prompts = "face"
|
||||
left_right_flip_augmentation = False
|
||||
|
||||
output_dir2 = os.path.join(out_root_dir, run_name + "_gene")
|
||||
input_dir2, n_imgs2, trigger_text2, segmentation_prompt2, captions2 = preprocess(
|
||||
output_dir2,
|
||||
concept_mode,
|
||||
input_zip_path=lora_training_urls,
|
||||
caption_text=caption_prefix,
|
||||
mask_target_prompts=mask_target_prompts,
|
||||
target_size=resolution,
|
||||
crop_based_on_salience=crop_based_on_salience,
|
||||
use_face_detection_instead=use_face_detection_instead,
|
||||
temp=clipseg_temperature,
|
||||
left_right_flip_augmentation=left_right_flip_augmentation,
|
||||
augment_imgs_up_to_n = augment_imgs_up_to_n,
|
||||
seed = seed,
|
||||
)
|
||||
|
||||
|
||||
# Merge the two preprocessing steps:
|
||||
n_imgs = n_imgs1 + n_imgs2
|
||||
captions = captions1 + captions2
|
||||
trigger_text = trigger_text1
|
||||
segmentation_prompt = segmentation_prompt1
|
||||
|
||||
# Create merged outdir:
|
||||
output_dir = os.path.join(out_root_dir, run_name + "_combined")
|
||||
input_dir = os.path.join(output_dir, "images_out")
|
||||
os.makedirs(input_dir, exist_ok=True)
|
||||
|
||||
# Merge the two preprocessed datasets:
|
||||
merge_datasets(input_dir1, input_dir2, input_dir, token_dict.keys())
|
||||
|
||||
else: # normal, single token run:
|
||||
|
||||
output_dir = os.path.join(out_root_dir, run_name)
|
||||
input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
|
||||
output_dir,
|
||||
concept_mode,
|
||||
input_zip_path=lora_training_urls,
|
||||
caption_text=caption_prefix,
|
||||
mask_target_prompts=mask_target_prompts,
|
||||
target_size=resolution,
|
||||
crop_based_on_salience=crop_based_on_salience,
|
||||
use_face_detection_instead=use_face_detection_instead,
|
||||
temp=clipseg_temperature,
|
||||
left_right_flip_augmentation=left_right_flip_augmentation,
|
||||
augment_imgs_up_to_n = augment_imgs_up_to_n,
|
||||
seed = seed,
|
||||
)
|
||||
|
||||
|
||||
if not debug:
|
||||
yield CogOutput(name=name, progress=0.05)
|
||||
|
||||
# Make a dict of all the arguments and save it to args.json:
|
||||
args_dict = {
|
||||
"name": name,
|
||||
"checkpoint": "juggernaut",
|
||||
"concept_mode": concept_mode,
|
||||
"input_images": str(lora_training_urls),
|
||||
"num_training_images": n_imgs,
|
||||
"num_augmented_images": len(captions),
|
||||
"seed": seed,
|
||||
"resolution": resolution,
|
||||
"train_batch_size": train_batch_size,
|
||||
"num_train_epochs": num_train_epochs,
|
||||
"max_train_steps": max_train_steps,
|
||||
"is_lora": is_lora,
|
||||
"prodigy_d_coef": prodigy_d_coef,
|
||||
"ti_lr": ti_lr,
|
||||
"ti_weight_decay": ti_weight_decay,
|
||||
"lora_weight_decay": lora_weight_decay,
|
||||
"l1_penalty": l1_penalty,
|
||||
"lora_param_scaler": lora_param_scaler,
|
||||
"lora_rank": lora_rank,
|
||||
"snr_gamma": snr_gamma,
|
||||
"trigger_text": trigger_text,
|
||||
"segmentation_prompt": segmentation_prompt,
|
||||
"crop_based_on_salience": crop_based_on_salience,
|
||||
"use_face_detection_instead": use_face_detection_instead,
|
||||
"clipseg_temperature": clipseg_temperature,
|
||||
"left_right_flip_augmentation": left_right_flip_augmentation,
|
||||
"augment_imgs_up_to_n": augment_imgs_up_to_n,
|
||||
"checkpointing_steps": checkpointing_steps,
|
||||
"run_name": run_name,
|
||||
"hard_pivot": hard_pivot,
|
||||
"off_ratio_power": off_ratio_power,
|
||||
"trainig_captions": captions[:50], # avoid sending back too many captions
|
||||
}
|
||||
|
||||
with open(os.path.join(output_dir, "training_args.json"), "w") as f:
|
||||
json.dump(args_dict, f, indent=4)
|
||||
|
||||
train_generator = main(
|
||||
pretrained_model,
|
||||
instance_data_dir=os.path.join(input_dir, "captions.csv"),
|
||||
output_dir=output_dir,
|
||||
config = TrainingConfig(
|
||||
name=name,
|
||||
lora_training_urls=lora_training_urls,
|
||||
concept_mode=concept_mode,
|
||||
sd_model_version=sd_model_version,
|
||||
seed=seed,
|
||||
resolution=resolution,
|
||||
train_batch_size=train_batch_size,
|
||||
num_train_epochs=num_train_epochs,
|
||||
max_train_steps=max_train_steps,
|
||||
gradient_accumulation_steps=gradient_accumulation_steps,
|
||||
l1_penalty=l1_penalty,
|
||||
prodigy_d_coef=prodigy_d_coef,
|
||||
checkpointing_steps=10000,
|
||||
ti_lr=ti_lr,
|
||||
ti_weight_decay=ti_weight_decay,
|
||||
snr_gamma=snr_gamma,
|
||||
lora_weight_decay=lora_weight_decay,
|
||||
token_dict=token_dict,
|
||||
inserting_list_tokens=all_token_lists,
|
||||
verbose=verbose,
|
||||
checkpointing_steps=checkpointing_steps,
|
||||
scale_lr=False,
|
||||
allow_tf32=True,
|
||||
mixed_precision="bf16",
|
||||
#mixed_precision="fp16", # this 100% breaks training... Figure out why!!?
|
||||
device="cuda:0",
|
||||
unet_lr=unet_lr,
|
||||
lora_rank=lora_rank,
|
||||
is_lora=is_lora,
|
||||
args_dict=args_dict,
|
||||
use_dora=use_dora,
|
||||
caption_model="blip",
|
||||
n_tokens=n_tokens,
|
||||
verbose=True,
|
||||
debug=debug,
|
||||
hard_pivot=hard_pivot,
|
||||
off_ratio_power=off_ratio_power,
|
||||
)
|
||||
|
||||
train_generator = train(config=config)
|
||||
print(f"Debug: {debug}")
|
||||
|
||||
while True:
|
||||
try:
|
||||
@@ -403,33 +146,9 @@ class Predictor(BasePredictor):
|
||||
if not debug:
|
||||
yield CogOutput(name=name, progress=np.round(progress_f, 2))
|
||||
except StopIteration as e:
|
||||
output_save_dir, validation_prompts = e.value # Capture the return value
|
||||
config, output_save_dir = e.value # Capture the return value
|
||||
break
|
||||
|
||||
if not debug:
|
||||
keys_to_keep = [
|
||||
"name",
|
||||
"checkpoint",
|
||||
"concept_mode",
|
||||
"input_images",
|
||||
"num_training_images",
|
||||
"seed",
|
||||
"resolution",
|
||||
"max_train_steps",
|
||||
"lora_rank",
|
||||
"trigger_text",
|
||||
"left_right_flip_augmentation",
|
||||
"run_name",
|
||||
"trainig_captions"]
|
||||
args_dict = {k: v for k, v in args_dict.items() if k in keys_to_keep}
|
||||
|
||||
args_dict["grid_prompts"] = validation_prompts
|
||||
|
||||
# save final training_args:
|
||||
final_args_dict_path = os.path.join(output_dir, "training_args.json")
|
||||
with open(final_args_dict_path, "w") as f:
|
||||
json.dump(args_dict, f, indent=4)
|
||||
|
||||
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
|
||||
out_path = f"{clean_filename(name)}_eden_concept_lora_{int(time.time())}.tar"
|
||||
directory = cogPath(output_save_dir)
|
||||
@@ -443,18 +162,18 @@ class Predictor(BasePredictor):
|
||||
|
||||
# Add instructions README:
|
||||
tar.add("instructions_README.md", arcname="README.md")
|
||||
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
|
||||
if sd_model_version == "sd15":
|
||||
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
|
||||
|
||||
attributes = {}
|
||||
attributes['grid_prompts'] = validation_prompts
|
||||
runtime = time.time() - start_time
|
||||
attributes['job_time_seconds'] = runtime
|
||||
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
||||
attributes['job_time_seconds'] = config.job_time
|
||||
|
||||
print(f"LORA training finished in {runtime:.1f} seconds")
|
||||
print(f"LORA training finished in {config.job_time:.1f} seconds")
|
||||
print(f"Returning {out_path}")
|
||||
|
||||
if DEBUG_MODE or debug:
|
||||
yield cogPath(out_path)
|
||||
else:
|
||||
# clear the output_directory to avoid running out of space on the machine:
|
||||
#shutil.rmtree(output_dir)
|
||||
yield CogOutput(files=[cogPath(out_path)], name=name, thumbnails=[cogPath(validation_grid_img_path)], attributes=args_dict, isFinal=True, progress=1.0)
|
||||
yield CogOutput(files=[cogPath(out_path)], name=name, thumbnails=[cogPath(validation_grid_img_path)], attributes=config.dict(), isFinal=True, progress=1.0)
|
||||
@@ -0,0 +1,22 @@
|
||||
torch==2.1.0
|
||||
torchaudio==2.1.0
|
||||
torchvision==0.16.0
|
||||
transformers==4.38.0
|
||||
diffusers==0.26.0
|
||||
tokenizers==0.15.2
|
||||
huggingface-hub==0.22.2
|
||||
ujson==5.10.0
|
||||
scipy==1.14.0
|
||||
peft==0.10.0
|
||||
invisible-watermark==0.2.0
|
||||
pandas==2.2.1
|
||||
numpy==1.26.4
|
||||
opencv-python==4.10.0.84
|
||||
mediapipe==0.10.14
|
||||
openai==1.35.13
|
||||
python-dotenv==1.0.1
|
||||
prodigyopt==1.0
|
||||
omegaconf==2.3.0
|
||||
ujson==5.10.0
|
||||
bitsandbytes==0.43.1
|
||||
setuptools==70.3.0
|
||||
@@ -0,0 +1,161 @@
|
||||
|
||||
"""
|
||||
Faces:
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
|
||||
|
||||
Objects:
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/koji_color.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/plantoid_imgs.zip
|
||||
|
||||
Styles:
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
|
||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
|
||||
|
||||
"""
|
||||
|
||||
import random, os, ast, json, shutil
|
||||
from itertools import product
|
||||
import time
|
||||
from tqdm import tqdm
|
||||
|
||||
random.seed(int(1000*time.time()))
|
||||
|
||||
def hamming_distance(dict1, dict2):
|
||||
distance = 0
|
||||
for key in dict1.keys():
|
||||
if dict1[key] != dict2.get(key, None):
|
||||
distance += 1
|
||||
return distance
|
||||
|
||||
#######################################################################################
|
||||
|
||||
# Setup the base experiment config:
|
||||
exp_name = "beeple"
|
||||
caption_prefix = ""
|
||||
mask_target_prompts = ""
|
||||
n_exp = 200 # how many random experiment settings to generate
|
||||
min_hamming_distance = 1 # min_n_params that have to be different from any previous experiment to be scheduled
|
||||
nohup = True
|
||||
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
|
||||
|
||||
# Define training hyperparameters and their possible values
|
||||
# The params are sampled stochastically, so if you want to use a specific value more often, just put it in multiple times
|
||||
|
||||
hyperparameters = {
|
||||
"output_dir": [f"lora_models/{exp_name}"],
|
||||
"sd_model_version": ["sdxl"],
|
||||
"lora_training_urls": [
|
||||
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
|
||||
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
|
||||
|
||||
],
|
||||
"concept_mode": ['style'],
|
||||
"sample_imgs_lora_scale": [0.8],
|
||||
"disable_ti": ['false', 'true'],
|
||||
"seed": [0],
|
||||
"resolution": [512],
|
||||
"train_batch_size": [4],
|
||||
"n_sample_imgs": [8],
|
||||
"max_train_steps": [1200],
|
||||
"checkpointing_steps": [200],
|
||||
"gradient_accumulation_steps": [1],
|
||||
|
||||
"n_tokens": [2],
|
||||
"ti_lr": [0.001],
|
||||
"ti_weight_decay": [0.001],
|
||||
"l1_penalty": [0.0],
|
||||
"token_warmup_steps": [0],
|
||||
"tok_cov_reg_w": [2000],
|
||||
|
||||
"unet_lr": [0.0002, 0.00005],
|
||||
"lora_alpha_multiplier": [1.0],
|
||||
"prodigy_d_coef": [1.0],
|
||||
"lora_weight_decay": [0.001],
|
||||
"lora_rank": [16],
|
||||
"use_dora": ['false'],
|
||||
|
||||
"unet_optimizer_type": ['AdamW8bit'],
|
||||
"is_lora": ['false'],
|
||||
|
||||
"text_encoder_lora_optimizer": [None],
|
||||
"text_encoder_lora_lr": [0.0e-4],
|
||||
|
||||
"snr_gamma": [5.0],
|
||||
"caption_model": ["blip", "gpt4-v"],
|
||||
"augment_imgs_up_to_n": [40],
|
||||
"verbose": ['true'],
|
||||
"debug": ['true']
|
||||
}
|
||||
|
||||
#######################################################################################
|
||||
|
||||
# Create a set to hold the combinations that have already been run
|
||||
scheduled_experiments = set()
|
||||
|
||||
# if config_output_dir exists, remove it:
|
||||
config_output_dir = f"gridsearch_configs/{exp_name}"
|
||||
shutil.rmtree(config_output_dir, ignore_errors=True)
|
||||
os.makedirs(config_output_dir, exist_ok=True)
|
||||
|
||||
# Open the shell script file
|
||||
try_sampling_n_times = 200
|
||||
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
|
||||
resamples, combination = 0, None
|
||||
|
||||
while resamples < try_sampling_n_times:
|
||||
experiment_settings = {name: random.choice(values) for name, values in hyperparameters.items()}
|
||||
resamples += 1
|
||||
|
||||
min_distance = float('inf')
|
||||
for str_experiment_settings in scheduled_experiments:
|
||||
existing_experiment_settings = dict(sorted(ast.literal_eval(str_experiment_settings)))
|
||||
distance = hamming_distance(experiment_settings, existing_experiment_settings)
|
||||
min_distance = min(min_distance, distance)
|
||||
|
||||
if min_distance >= min_hamming_distance:
|
||||
str_experiment_settings = str(sorted(experiment_settings.items()))
|
||||
scheduled_experiments.add(str_experiment_settings)
|
||||
# Save the experiment to a JSON file
|
||||
config_filename = f"{config_output_dir}/{exp_name}_{exp_index:03d}.json"
|
||||
dirname = os.path.dirname(config_filename)
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
|
||||
# Make some final adjustments to the experiment settings before saving to disk:
|
||||
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
|
||||
|
||||
with open(config_filename, "w") as f:
|
||||
json.dump(experiment_settings, f, indent=4)
|
||||
break
|
||||
|
||||
if resamples >= try_sampling_n_times:
|
||||
print(f"\nCould not find a new experiment_setting after random sampling {try_sampling_n_times} times, dumping all experiment_settings to .json files")
|
||||
break
|
||||
|
||||
print(f"\n\n---> Saved {len(scheduled_experiments)} experiment configurations to {config_output_dir}")
|
||||
|
||||
def generate_sh_script(folder_path, output_sh_path):
|
||||
# Get a list of JSON files in the folder, sorted alphabetically
|
||||
json_files = sorted([f for f in os.listdir(folder_path) if f.endswith('.json')])
|
||||
|
||||
# Open the output .sh file for writing
|
||||
with open(output_sh_path, 'w') as sh_file:
|
||||
# Write the shebang line for a bash script
|
||||
sh_file.write("#!/bin/bash\n\n")
|
||||
|
||||
# Write a command for each JSON file
|
||||
for json_file in json_files:
|
||||
file_path = os.path.join("scripts/", folder_path, json_file)
|
||||
command = f"python main.py {file_path}\n"
|
||||
|
||||
if nohup:
|
||||
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
|
||||
|
||||
sh_file.write(command)
|
||||
|
||||
generate_sh_script(config_output_dir, output_sh_path)
|
||||
print(f"\n---> Saved the executable shell script to {output_sh_path}")
|
||||
Executable
+236
@@ -0,0 +1,236 @@
|
||||
import argparse
|
||||
from trainer.inference import render_images_eval
|
||||
from trainer.utils.json_stuff import save_as_json
|
||||
from trainer.config import TrainingConfig
|
||||
from trainer.models import pretrained_models
|
||||
from trainer.utils.io import download
|
||||
import clip
|
||||
from PIL import Image
|
||||
import torch
|
||||
import numpy as np
|
||||
import os
|
||||
from creator_lora.models.resnet50 import ResNet50MLP
|
||||
|
||||
"""
|
||||
todos:
|
||||
- run eval on user-defined captions
|
||||
"""
|
||||
|
||||
aesthetic_model_checkpoint_filename = "aesthetic_score_best_model.pth"
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
def get_filenames_in_a_folder(folder: str):
|
||||
"""
|
||||
returns the list of paths to all the files in a given folder
|
||||
"""
|
||||
|
||||
if folder[-1] == '/':
|
||||
folder = folder[:-1]
|
||||
|
||||
files = os.listdir(folder)
|
||||
files = [f'{folder}/' + x for x in files]
|
||||
return files
|
||||
|
||||
def get_all_jpg_filenames(folder):
|
||||
all_filenames = get_filenames_in_a_folder(folder=folder)
|
||||
jpg_filenames = [filename for filename in all_filenames if filename.lower().endswith('.jpg')]
|
||||
assert len(jpg_filenames)>0, f"Expected to find at least 1 jpg file but got 0"
|
||||
return jpg_filenames
|
||||
|
||||
def filter_prompt(prompt, remove_this = "in the style of <s0><s1>,", replace_with = ""):
|
||||
assert remove_this in prompt, f"Expected '{remove_this}' to be present in the prompt: '{prompt}'"
|
||||
return prompt.replace(
|
||||
remove_this,
|
||||
replace_with
|
||||
)
|
||||
|
||||
def get_similarity_matrix(a, b, eps=1e-8):
|
||||
"""
|
||||
finds the cosine similarity matrix between each item of a w.r.t each item of b
|
||||
a and b are expected to be 2 dimensional
|
||||
added eps for numerical stability
|
||||
source: https://stackoverflow.com/a/58144658
|
||||
"""
|
||||
a_n, b_n = a.norm(dim=1)[:, None], b.norm(dim=1)[:, None]
|
||||
a_norm = a / torch.max(a_n, eps * torch.ones_like(a_n))
|
||||
b_norm = b / torch.max(b_n, eps * torch.ones_like(b_n))
|
||||
sim_mt = torch.mm(a_norm, b_norm.transpose(0, 1))
|
||||
return sim_mt
|
||||
|
||||
class Evaluation:
|
||||
def __init__(self, image_filenames: list):
|
||||
self.image_filenames = image_filenames
|
||||
self.image_features = None
|
||||
|
||||
def obtain_image_features(self):
|
||||
|
||||
if self.image_features is None:
|
||||
all_image_features = []
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
|
||||
for f in self.image_filenames:
|
||||
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
|
||||
with torch.no_grad():
|
||||
image_features = model.encode_image(image)
|
||||
all_image_features.append(image_features.float())
|
||||
|
||||
all_image_features = torch.cat(all_image_features, dim = 0)
|
||||
self.image_features = all_image_features
|
||||
|
||||
return self.image_features
|
||||
|
||||
def obtain_text_features(self, prompts: list, device):
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
text = clip.tokenize(prompts).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
text_features = model.encode_text(text)
|
||||
return text_features
|
||||
|
||||
def training_image_alignment(self, device, training_image_filenames: list):
|
||||
generated_image_features = self.obtain_image_features()
|
||||
|
||||
training_image_features = []
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
|
||||
for f in training_image_filenames:
|
||||
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
|
||||
with torch.no_grad():
|
||||
image_features = model.encode_image(image)
|
||||
training_image_features.append(image_features.float())
|
||||
|
||||
training_image_features = torch.cat(training_image_features, dim = 0)
|
||||
return get_similarity_matrix(a=generated_image_features, b=training_image_features).mean().item()
|
||||
|
||||
|
||||
def image_text_alignment(self, device, prompts: list):
|
||||
|
||||
image_features = self.obtain_image_features().to(device)
|
||||
assert image_features.shape[0] == len(prompts), f'Expected len(prompts) ({len(prompts)}) to have the same number of prompts as the number of images provided: {image_features.shape}'
|
||||
text_features = self.obtain_text_features(prompts=prompts, device=device)
|
||||
cossim = torch.nn.functional.cosine_similarity(
|
||||
text_features, image_features, dim = -1
|
||||
).mean().item()
|
||||
return cossim
|
||||
|
||||
def clip_diversity(self, device: str):
|
||||
"""
|
||||
higher = more diverse
|
||||
"""
|
||||
all_image_features = self.obtain_image_features().to(device)
|
||||
|
||||
distances = 1 - get_similarity_matrix(all_image_features, all_image_features)
|
||||
assert distances.shape == (
|
||||
all_image_features.shape[0],
|
||||
all_image_features.shape[0]
|
||||
), f'Expected the shape of the distance matrix to be (num_images, num_images) i.e {(all_image_features.shape[0], all_image_features.shape[0])} but got: {distances.shape}'
|
||||
distances = distances.detach().cpu().numpy()
|
||||
# Get the upper triangle:
|
||||
upper_triangle = np.triu(distances, k=1).flatten()
|
||||
return upper_triangle.mean().item()
|
||||
|
||||
def aesthetic_score(self, device: str, checkpoint_path: str):
|
||||
# assert os.path.exists(checkpoint_path), f"invalid checkpoint_path: {checkpoint_path}"
|
||||
model = ResNet50MLP(
|
||||
model_path=checkpoint_path,
|
||||
device = device
|
||||
)
|
||||
|
||||
scores = []
|
||||
for f in self.image_filenames:
|
||||
score = model.predict_score(pil_image=Image.open(f))
|
||||
scores.append(score)
|
||||
|
||||
return sum(scores)/len(scores)
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description="Script for generating images based on prompts and computing similarities.")
|
||||
|
||||
|
||||
parser.add_argument("--config_filename", type=str, required=True, default = "sdxl", help="path to config json file")
|
||||
parser.add_argument("--checkpoint_folder", type=str, required=True,
|
||||
help="Path to folder containing the checkpoint. Usually a folder which is named like: .../checkpoint-500")
|
||||
parser.add_argument("--output_json", type=str, required=True,
|
||||
help="Path to json where we save result values")
|
||||
parser.add_argument("--output_folder", type=str, required=True,
|
||||
help="style or face")
|
||||
parser.add_argument("--training_images_folder", type=str, required=True,
|
||||
help="path to folder containing training image jpg files. Usually the `images_in` folder")
|
||||
args = parser.parse_args()
|
||||
|
||||
## validate args
|
||||
assert os.path.exists(args.checkpoint_folder), f"Invalid lora_path: {args.checkpoint_folder}"
|
||||
assert os.path.exists(args.config_filename), f"Invalid lora_path: {args.config_filename}"
|
||||
assert os.path.exists(args.training_images_folder), f"Invalid training_images_folder: {args.training_images_folder}"
|
||||
return args
|
||||
|
||||
args = parse_arguments()
|
||||
|
||||
os.system(f"mkdir -p {args.output_folder}")
|
||||
if not os.path.exists(aesthetic_model_checkpoint_filename):
|
||||
download(
|
||||
url="https://edenartlab-lfs.s3.amazonaws.com/models/aesthetic_score_best_model.pth",
|
||||
folder="./",
|
||||
filepath=None
|
||||
)
|
||||
|
||||
config = TrainingConfig.from_json(args.config_filename)
|
||||
|
||||
image_filenames, prompts = render_images_eval(
|
||||
output_folder=args.output_folder,
|
||||
concept_mode=config.concept_mode,
|
||||
render_size=(1024,1024),
|
||||
checkpoint_folder=args.checkpoint_folder,
|
||||
pretrained_model=pretrained_models[config.sd_model_version],
|
||||
seed=0,
|
||||
is_lora = config.is_lora,
|
||||
trigger_text='TOK' if config.concept_mode != "style" else ", in the style of TOK"
|
||||
)
|
||||
|
||||
|
||||
print(f"Eval prompts:")
|
||||
for i, p in enumerate(prompts):
|
||||
print(f"{i}:{p}")
|
||||
|
||||
eval = Evaluation(image_filenames=image_filenames)
|
||||
clip_diversity = eval.clip_diversity(device=device)
|
||||
|
||||
|
||||
aesthetic_score = eval.aesthetic_score(device=device, checkpoint_path=aesthetic_model_checkpoint_filename)
|
||||
image_text_alignment = eval.image_text_alignment(device=device, prompts=prompts)
|
||||
training_image_alignment = eval.training_image_alignment(
|
||||
device=device,
|
||||
training_image_filenames=get_all_jpg_filenames(folder=args.training_images_folder)
|
||||
)
|
||||
|
||||
result = {
|
||||
"sd_model_version": config.sd_model_version,
|
||||
"checkpoint_folder": os.path.abspath(args.checkpoint_folder),
|
||||
"concept_mode": config.concept_mode,
|
||||
"output_folder": args.output_folder,
|
||||
"training_images_folder":args.training_images_folder,
|
||||
"scores": {
|
||||
"clip_diversity": clip_diversity,
|
||||
"aesthetic_score": aesthetic_score,
|
||||
"image_text_alignment": image_text_alignment,
|
||||
"training_image_alignment": training_image_alignment
|
||||
}
|
||||
}
|
||||
|
||||
save_as_json(
|
||||
dictionary_or_list=result,
|
||||
filename=args.output_json
|
||||
)
|
||||
print(f"Eval complete. Saved results here: {args.output_json}")
|
||||
|
||||
"""
|
||||
Example command:
|
||||
|
||||
python3 evaluate.py \
|
||||
--output_folder eval_images \
|
||||
--checkpoint_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0 \
|
||||
--output_json eval_results_style.json \
|
||||
--config_filename lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0/training_args.json \
|
||||
--training_images_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/images_in
|
||||
"""
|
||||
@@ -0,0 +1,142 @@
|
||||
import os
|
||||
import json
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from sklearn.linear_model import LinearRegression
|
||||
from sklearn.metrics import r2_score
|
||||
|
||||
|
||||
# Define paths
|
||||
render_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/lora_models"
|
||||
config_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/xander_adiff_lora"
|
||||
|
||||
ignore_threshold_relative = 0.0 # ignore any datapoint with a score below this threshold
|
||||
|
||||
filters = {
|
||||
"resolution": 512
|
||||
}
|
||||
|
||||
output_dir = f"gridsearch_configs/results/{os.path.basename(config_dir)}"
|
||||
output_suffix = f"{os.path.basename(render_dir)}"
|
||||
|
||||
# Initialize a dictionary to hold parameter values and associated scores
|
||||
parameters = defaultdict(lambda: defaultdict(list))
|
||||
|
||||
# Step 1: Loop over each experiment subdirectory
|
||||
for i, exp_subdir in enumerate(sorted(os.listdir(render_dir))):
|
||||
exp_path = os.path.join(render_dir, exp_subdir)
|
||||
checkpoints_path = os.path.join(exp_path, "checkpoints")
|
||||
|
||||
# Step 2: Get the score by counting the number of .jpg files in the checkpoints subdir
|
||||
if os.path.isdir(checkpoints_path):
|
||||
score = sum(1 for _ in os.listdir(checkpoints_path) if _.endswith('.jpg'))
|
||||
|
||||
# Match the experiment folder with its corresponding JSON file
|
||||
json_file_name = exp_subdir.split('--')[0] + ".json"
|
||||
json_file_name = json_file_name.replace('__','_')
|
||||
json_path = os.path.join(config_dir, json_file_name)
|
||||
|
||||
# Step 3: Load the corresponding .json file
|
||||
if os.path.isfile(json_path):
|
||||
with open(json_path, 'r') as file:
|
||||
config = json.load(file)
|
||||
|
||||
# Filter out experiments that do not match the filters
|
||||
if not all(config[key] == value for key, value in filters.items()):
|
||||
continue
|
||||
|
||||
# Step 4: Append all key/value pairs to the total experiment dictionary
|
||||
for key, value in config.items():
|
||||
parameters[key]['values'].append(value)
|
||||
parameters[key]['scores'].append(score)
|
||||
else:
|
||||
print(f"Could not find JSON file for experiment {exp_subdir}")
|
||||
|
||||
|
||||
# Print the parameters['output_dir'] with the highest scores (there are usually multiple ties):
|
||||
max_score = max(parameters['output_dir']['scores'])
|
||||
best_output_dirs = [output_dir for output_dir, score in zip(parameters['output_dir']['values'], parameters['output_dir']['scores']) if score == max_score]
|
||||
for best_output_dir in best_output_dirs:
|
||||
print(f"Best output_dir: {best_output_dir} with score {max_score}")
|
||||
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from sklearn.linear_model import LinearRegression
|
||||
from sklearn.metrics import r2_score
|
||||
from sklearn.preprocessing import LabelEncoder
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
print(f"Saving results to {output_dir}...")
|
||||
|
||||
def plot_parameters(parameters):
|
||||
for param, data in parameters.items():
|
||||
values = np.array(data['values'])
|
||||
scores = np.array(data['scores'])
|
||||
|
||||
# filter based on the ignore_threshold:
|
||||
ignore_threshold = ignore_threshold_relative * np.max(scores)
|
||||
mask = scores > ignore_threshold
|
||||
values = values[mask]
|
||||
scores = scores[mask]
|
||||
|
||||
noise_strength_values = 0.02
|
||||
noise_strength_scores = 0.02
|
||||
|
||||
# Initialize variables for original categorical labels
|
||||
original_labels = None
|
||||
|
||||
# Determine if values are numeric
|
||||
if values.dtype.kind in 'bifc': # Numeric types
|
||||
# Add noise directly to values
|
||||
jittered_values = values + np.random.normal(0, noise_strength_values * (np.max(values) - np.min(values)), values.shape)
|
||||
else:
|
||||
# Encode string values to integers for plotting
|
||||
encoder = LabelEncoder()
|
||||
original_labels = values.copy()
|
||||
values = encoder.fit_transform(values)
|
||||
jittered_values = values + np.random.normal(0, 0.1, values.shape)
|
||||
|
||||
# Skip plotting if there is only one unique value for the parameter
|
||||
if len(np.unique(values)) <= 1:
|
||||
continue
|
||||
|
||||
# Fit a linear regression model to the encoded values if categorical
|
||||
model = LinearRegression()
|
||||
values_reshaped = values.reshape(-1, 1) # Reshape for sklearn
|
||||
model.fit(values_reshaped, scores)
|
||||
predicted_scores = model.predict(values_reshaped)
|
||||
|
||||
# add some jitter to the scores:
|
||||
jittered_scores = scores + np.random.normal(0, noise_strength_scores * np.max(scores), scores.shape)
|
||||
|
||||
# Calculate R² value
|
||||
r_squared = r2_score(scores, predicted_scores)
|
||||
|
||||
# Plot data points
|
||||
sns.scatterplot(x=jittered_values, y=jittered_scores, alpha=0.6)
|
||||
|
||||
# Plot trendline
|
||||
sns.lineplot(x=np.sort(values), y=predicted_scores[np.argsort(values)], color='red', label=f'R²={r_squared:.2f}')
|
||||
|
||||
# Set plot title and labels
|
||||
plt.title(f'Influence of {param} on the score')
|
||||
if original_labels is not None:
|
||||
# Set x-axis labels to the original categorical labels
|
||||
unique_values = np.unique(values)
|
||||
plt.xticks(ticks=unique_values, labels=encoder.inverse_transform(unique_values), rotation=45, ha='right')
|
||||
else:
|
||||
plt.xlabel(param)
|
||||
plt.ylabel('Score')
|
||||
plt.legend()
|
||||
|
||||
# Save and close the plot
|
||||
plt.savefig(f'{output_dir}/res_{param}_{output_suffix}.png')
|
||||
plt.close()
|
||||
|
||||
# Call the updated function with your parameters dictionary
|
||||
plot_parameters(parameters)
|
||||
@@ -0,0 +1,92 @@
|
||||
from diffusers import DDPMScheduler, EulerDiscreteScheduler, StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from peft import PeftModel
|
||||
import numpy as np
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
import os, json, random, time, sys
|
||||
|
||||
sys.path.append('.')
|
||||
sys.path.append('..')
|
||||
from trainer.models import load_models, pretrained_models
|
||||
from trainer.utils.val_prompts import val_prompts
|
||||
from trainer.utils.io import make_validation_img_grid
|
||||
from trainer.utils.utils import seed_everything, pick_best_gpu_id
|
||||
from trainer.inference import encode_prompt_advanced
|
||||
from trainer.checkpoint import load_checkpoint
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
model_version = "sd15"
|
||||
lora_path = 'lora_models/XANDER_SD15_SWEEP/sd15_face_sweep__004--29_20-43-17-sd15_face_dora_640_1.0_blip_800/checkpoints/checkpoint-800'
|
||||
lora_scales = np.linspace(0.6, 0.9, 4)
|
||||
token_scale = None # None means it well get automatically set using lora_scale
|
||||
render_size = (576, 704) # H,W
|
||||
n_imgs = 14
|
||||
n_loops = 2
|
||||
|
||||
n_steps = 35
|
||||
guidance_scale = 7.5
|
||||
seed = 12
|
||||
use_lightning = 0
|
||||
|
||||
#####################################################################################
|
||||
|
||||
pretrained_model = pretrained_models[model_version]
|
||||
output_dir = f'rendered_images/{lora_path.split("/")[-1]}'
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
seed_everything(seed)
|
||||
pick_best_gpu_id()
|
||||
|
||||
pipe = load_checkpoint(
|
||||
pretrained_model_version=model_version,
|
||||
pretrained_model_path=pretrained_model["path"],
|
||||
checkpoint_folder=lora_path,
|
||||
is_lora=True,
|
||||
device="cuda:0"
|
||||
)
|
||||
|
||||
if use_lightning:
|
||||
repo = "ByteDance/SDXL-Lightning"
|
||||
ckpt = "sdxl_lightning_8step_lora.safetensors" # Use the correct ckpt for your step setting!
|
||||
pipe.load_lora_weights(hf_hub_download(repo, ckpt))
|
||||
pipe.fuse_lora()
|
||||
n_steps = 8
|
||||
guidance_scale=1.5
|
||||
|
||||
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
|
||||
training_args = json.load(f)
|
||||
|
||||
if training_args["concept_mode"] == "style":
|
||||
validation_prompts_raw = random.choices(val_prompts['style'], k=n_imgs)
|
||||
elif training_args["concept_mode"] == "face":
|
||||
validation_prompts_raw = random.choices(val_prompts['face'], k=n_imgs)
|
||||
else:
|
||||
validation_prompts_raw = random.choices(val_prompts['object'], k=n_imgs)
|
||||
|
||||
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
|
||||
pipeline_args = {
|
||||
"num_inference_steps": n_steps,
|
||||
"guidance_scale": guidance_scale,
|
||||
"height": render_size[0],
|
||||
"width": render_size[1],
|
||||
}
|
||||
for jj in range(n_loops):
|
||||
for i in range(len(validation_prompts_raw)):
|
||||
for lora_scale in lora_scales:
|
||||
seed += 1
|
||||
pipe = set_adapter_scales(pipe, lora_scale=lora_scale)
|
||||
generator = torch.Generator(device='cuda').manual_seed(seed)
|
||||
|
||||
c, uc, pc, puc = encode_prompt_advanced(pipe, lora_path, validation_prompts_raw[i], negative_prompt, lora_scale, guidance_scale, concept_mode = training_args["concept_mode"], token_scale = token_scale)
|
||||
|
||||
pipeline_args['prompt_embeds'] = c
|
||||
pipeline_args['negative_prompt_embeds'] = uc
|
||||
if pretrained_model['version'] == 'sdxl':
|
||||
pipeline_args['pooled_prompt_embeds'] = pc
|
||||
pipeline_args['negative_pooled_prompt_embeds'] = puc
|
||||
|
||||
image = pipe(**pipeline_args, generator=generator).images[0]
|
||||
image.save(os.path.join(output_dir, f"{validation_prompts_raw[i][:40]}_seed_{seed}_{i}_lora_scale_{lora_scale:.2f}_{int(time.time())}.jpg"), format="JPEG", quality=95)
|
||||
|
||||
seed += 1
|
||||
@@ -1,26 +0,0 @@
|
||||
# Set GPU ID to run these jobs on:
|
||||
GPU_ID="device=3"
|
||||
|
||||
cog predict --gpus $GPU_ID \
|
||||
-i run_name="clipx_sdxl" \
|
||||
-i caption_prefix="in the style of TOK, " \
|
||||
-i concept_mode="style" \
|
||||
-i train_batch_size="4" \
|
||||
-i sd_model_version="sdxl" \
|
||||
-i max_train_steps="500" \
|
||||
-i checkpointing_steps="125" \
|
||||
-i debug="True" \
|
||||
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
|
||||
-i seed="0"
|
||||
|
||||
cog predict --gpus $GPU_ID \
|
||||
-i run_name="clipx_sd15" \
|
||||
-i caption_prefix="in the style of TOK, " \
|
||||
-i concept_mode="style" \
|
||||
-i train_batch_size="4" \
|
||||
-i sd_model_version="sd15" \
|
||||
-i max_train_steps="500" \
|
||||
-i checkpointing_steps="125" \
|
||||
-i debug="True" \
|
||||
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
|
||||
-i seed="0"
|
||||
@@ -0,0 +1,8 @@
|
||||
# Set GPU ID to run these jobs on:
|
||||
GPU_ID="device=0"
|
||||
|
||||
python main.py train_configs/training_args_face_sdxl.json
|
||||
python main.py train_configs/training_args_face_sd15.json
|
||||
python main.py train_configs/training_args_object.json
|
||||
python main.py train_configs/training_args_style_sd15.json
|
||||
python main.py train_configs/training_args_style_sdxl.json
|
||||
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"name": "xander_test",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
|
||||
"concept_mode": "face",
|
||||
"seed": 1,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 4,
|
||||
"max_train_steps": 200,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 100,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"disable_ti": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"name": "xander_sdxl",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
|
||||
"concept_mode": "face",
|
||||
"seed": 1,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 400,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 200,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"disable_ti": false,
|
||||
"n_tokens": 2,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.00,
|
||||
"lora_rank": 4,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"name": "xander_sd15",
|
||||
"sd_model_version": "sd15",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
|
||||
"concept_mode": "face",
|
||||
"seed": 0,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 600,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 300,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"remove_ti_token_from_prompts": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"name": "xander_sdxl",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
|
||||
"concept_mode": "face",
|
||||
"seed": 1,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 400,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 200,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"disable_ti": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"name": "banny_sd15",
|
||||
"sd_model_version": "sd15",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny.zip",
|
||||
"concept_mode": "face",
|
||||
"sample_imgs_lora_scale": 0.8,
|
||||
"seed": 0,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 800,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 200,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"remove_ti_token_from_prompts": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
{
|
||||
"name": "clipx_sd15",
|
||||
"sd_model_version": "sd15",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
|
||||
"concept_mode": "style",
|
||||
"seed": 0,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 400,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 200,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"remove_ti_token_from_prompts": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"name": "clipx_sdxl",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
|
||||
"concept_mode": "style",
|
||||
"sample_imgs_lora_scale": 0.7,
|
||||
"seed": 1,
|
||||
"resolution": 512,
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"max_train_steps": 400,
|
||||
"token_warmup_steps": 0,
|
||||
"checkpointing_steps": 200,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0005,
|
||||
|
||||
"remove_ti_token_from_prompts": false,
|
||||
"text_encoder_lora_optimizer": null,
|
||||
"text_encoder_lora_lr": 1.0e-4,
|
||||
"text_encoder_lora_weight_decay": 1e-5,
|
||||
"text_encoder_lora_rank": 12,
|
||||
|
||||
"unet_lr": 0.001,
|
||||
"lora_rank": 16,
|
||||
"use_dora": false,
|
||||
"caption_model": "blip",
|
||||
"debug": true
|
||||
}
|
||||
@@ -0,0 +1,276 @@
|
||||
import os, json
|
||||
from peft.utils import get_peft_model_state_dict
|
||||
from diffusers.utils import (
|
||||
convert_all_state_dict_to_peft,
|
||||
convert_state_dict_to_diffusers,
|
||||
convert_state_dict_to_kohya,
|
||||
convert_unet_state_dict_to_peft,
|
||||
)
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from safetensors.torch import load_file, save_file
|
||||
import torch
|
||||
from diffusers import EulerDiscreteScheduler
|
||||
from peft import PeftModel
|
||||
from .utils.json_stuff import save_as_json
|
||||
|
||||
from typing import Dict
|
||||
from trainer.embedding_handler import TokenEmbeddingsHandler
|
||||
|
||||
def load_ti_embeddings(pipe, save_path):
|
||||
# Load the textual_inversion token embeddings into the pipeline:
|
||||
try: #SDXL
|
||||
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
|
||||
except: #SD15
|
||||
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
|
||||
|
||||
embeddings_path = [f for f in os.listdir(save_path) if f.endswith("embeddings.safetensors")][0]
|
||||
print(f"Loading pretrained token embeddings from {embeddings_path}")
|
||||
handler.load_embeddings(os.path.join(save_path, embeddings_path))
|
||||
|
||||
|
||||
def set_adapter_scales(pipe, lora_scale = 1.0):
|
||||
"""
|
||||
update the pipe with the lora model and the token embeddings
|
||||
"""
|
||||
|
||||
# this loads the lora model into the pipeline at full strength (1.0)
|
||||
#pipe.unet.load_adapter(lora_path, "eden_lora")
|
||||
#peft_model.set_adapter(["adapter1", "adapter2"]) # activate both adapters
|
||||
|
||||
# First lets see if any lora's are active and unload them:
|
||||
#pipe.unet.unmerge_adapter()
|
||||
|
||||
list_adapters_component_wise = pipe.get_list_adapters()
|
||||
print(f"list_adapters_component_wise: {list_adapters_component_wise}")
|
||||
|
||||
if 1:
|
||||
for key in list_adapters_component_wise:
|
||||
adapter_names = list_adapters_component_wise[key]
|
||||
for adapter_name in adapter_names:
|
||||
print(f"Set adapter '{adapter_name}' of '{key}' with scale = {lora_scale:.2f}")
|
||||
pipe.set_adapters(adapter_name, adapter_weights=[lora_scale])
|
||||
|
||||
#pipe.unet.merge_adapter()
|
||||
|
||||
return pipe
|
||||
|
||||
|
||||
def remove_delimiter_characters(name: str):
|
||||
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
|
||||
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
|
||||
|
||||
# Convert to WebUI format
|
||||
def convert_pytorch_lora_safetensors_to_webui(
|
||||
pytorch_lora_weights_filename: str,
|
||||
output_filename: str
|
||||
):
|
||||
assert os.path.exists(pytorch_lora_weights_filename), f"Invalid path: {pytorch_lora_weights_filename}"
|
||||
lora_state_dict = load_file(pytorch_lora_weights_filename)
|
||||
peft_state_dict = convert_all_state_dict_to_peft(lora_state_dict)
|
||||
kohya_state_dict = convert_state_dict_to_kohya(peft_state_dict)
|
||||
|
||||
# This is a very custom hack because for some reason these 'base_model_model_' prefixes are added to the keys and ComfyUI does not like them...
|
||||
replace_dict = {"base_model_model_": ""}
|
||||
# enumerate and apply replace_dict:
|
||||
for key in list(kohya_state_dict.keys()):
|
||||
for old_key, new_key in replace_dict.items():
|
||||
if old_key in key:
|
||||
new_key = key.replace(old_key, new_key)
|
||||
kohya_state_dict[new_key] = kohya_state_dict.pop(key)
|
||||
|
||||
save_file(kohya_state_dict, output_filename)
|
||||
|
||||
def save_checkpoint(
|
||||
output_dir: str,
|
||||
global_step: int,
|
||||
unet,
|
||||
embedding_handler,
|
||||
token_dict: dict,
|
||||
is_lora: bool,
|
||||
unet_lora_parameters,
|
||||
pretrained_model_version: str,
|
||||
name: str = None,
|
||||
text_encoder_peft_models: list = [None]
|
||||
):
|
||||
"""
|
||||
Save the model's embeddings and special parameters (Lora) to the specified directory.
|
||||
|
||||
Note: This function directly corresponds to the `load_checkpoint` method
|
||||
|
||||
Args:
|
||||
`output_dir` (str): The directory path where the checkpoint will be saved.
|
||||
`global_step` (int): The current global step or epoch number.
|
||||
`unet`: The main model to save.
|
||||
`embedding_handler`: The handler for saving embeddings.
|
||||
`token_dict` (dict): Special parameters associated with the model.
|
||||
`is_lora` (bool): Whether the model includes LoRA components.
|
||||
`unet_lora_parameters`: Parameters associated with the LoRA components.
|
||||
`name` (str, optional): Name identifier for the checkpoint. Defaults to None.
|
||||
`text_encoder_peft_models` (list, optional): List of additional text encoder models to save. Defaults to None.
|
||||
|
||||
Returns:
|
||||
None
|
||||
|
||||
Saves:
|
||||
- {name}_embeddings.safetensors: Embeddings of the model.
|
||||
- special_params.json: Special parameters of the model.
|
||||
|
||||
If `text_encoder_peft_models` is provided, saves each model in a separate directory with the
|
||||
following structure:
|
||||
- text_encoder_lora_{index}/
|
||||
- adapter_config.json
|
||||
- adapter_model.safetensors
|
||||
- README.md
|
||||
|
||||
If `is_lora` is True, saves additional LoRA-related data:
|
||||
- LoRA weights
|
||||
- LoRA weights converted for web UI
|
||||
|
||||
If `is_lora` is False then it assumes that it's a vanilla unet model and saves it in the usual huggingface way.
|
||||
|
||||
"""
|
||||
print(f"Saving checkpoint at step.. {global_step}")
|
||||
name = remove_delimiter_characters(name)
|
||||
|
||||
embedding_handler.save_embeddings(
|
||||
os.path.join(
|
||||
output_dir,
|
||||
f"{name}_{pretrained_model_version}_embeddings.safetensors"
|
||||
)
|
||||
)
|
||||
|
||||
save_as_json(
|
||||
token_dict,
|
||||
filename = os.path.join(
|
||||
output_dir, "special_params.json"
|
||||
)
|
||||
)
|
||||
|
||||
if is_lora:
|
||||
assert len(unet_lora_parameters) > 0, f"Expected len(unet_lora_parameters) to be greater than zero if is_lora is True"
|
||||
|
||||
# This saves adapter_config.json:
|
||||
# TODO: adjust inference.py so it can load everything without needing this file
|
||||
unet.save_pretrained(save_directory = output_dir)
|
||||
|
||||
text_encoder_lora_layers = [None, None]
|
||||
for idx, model in enumerate(text_encoder_peft_models):
|
||||
if model is not None:
|
||||
lora_tensors = get_peft_model_state_dict(model)
|
||||
text_encoder_lora_layers[idx] = convert_state_dict_to_diffusers(lora_tensors)
|
||||
|
||||
lora_tensors = get_peft_model_state_dict(unet)
|
||||
unet_lora_layers_to_save = convert_state_dict_to_diffusers(lora_tensors)
|
||||
|
||||
if pretrained_model_version == "sdxl":
|
||||
print("Saving LoRA weights for SDXL model...")
|
||||
StableDiffusionXLPipeline.save_lora_weights(
|
||||
output_dir,
|
||||
unet_lora_layers=unet_lora_layers_to_save,
|
||||
text_encoder_lora_layers=text_encoder_lora_layers[0],
|
||||
text_encoder_2_lora_layers=text_encoder_lora_layers[1],
|
||||
)
|
||||
elif pretrained_model_version == "sd15":
|
||||
print("Saving LoRA weights for SD15 model...")
|
||||
StableDiffusionPipeline.save_lora_weights(
|
||||
output_dir,
|
||||
unet_lora_layers=unet_lora_layers_to_save,
|
||||
text_encoder_lora_layers=text_encoder_lora_layers[0],
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid pretrained_model_version: {pretrained_model_version}. Expected one of: 'sdxl' or 'sd15'"
|
||||
)
|
||||
|
||||
convert_pytorch_lora_safetensors_to_webui(
|
||||
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
|
||||
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_LoRa.safetensors")
|
||||
)
|
||||
else:
|
||||
# Save the entire, finetuned unet weights:
|
||||
unet.save_pretrained(save_directory = output_dir)
|
||||
|
||||
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
|
||||
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
|
||||
for file in to_remove:
|
||||
file_path = os.path.join(output_dir, file)
|
||||
if os.path.exists(file_path):
|
||||
os.remove(file_path)
|
||||
|
||||
return
|
||||
|
||||
def load_checkpoint(
|
||||
pretrained_model_version: str,
|
||||
pretrained_model_path: str,
|
||||
lora_save_path: str,
|
||||
is_lora: bool,
|
||||
device: str,
|
||||
lora_scale: float = 1.0,
|
||||
):
|
||||
"""
|
||||
Load a pre-trained model checkpoint and prepare it for inference.
|
||||
|
||||
Note: This function directly corresponds to the `save_checkpoint` method
|
||||
|
||||
Args:
|
||||
`pretrained_model_version` (`str`): Version of the pre-trained model (`sd15` or `sdxl`).
|
||||
`pretrained_model_path` (`str`): Path to the pre-trained model file.
|
||||
`lora_save_path` (`str`): Path to the LoRa checkpoint folder.
|
||||
`is_lora` (`bool`): Whether LoRA model components are used.
|
||||
`device` (Union[`str`, `torch.device`]): Device for inference.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If an unsupported `pretrained_model_version` is provided.
|
||||
"""
|
||||
|
||||
assert os.path.exists(pretrained_model_path), f"Invalid pretrained_model_path: {pretrained_model_path}"
|
||||
|
||||
if pretrained_model_version == "sd15":
|
||||
pipe = StableDiffusionPipeline.from_single_file(
|
||||
pretrained_model_path, torch_dtype=torch.float16, use_safetensors=True)
|
||||
elif pretrained_model_version == "sdxl":
|
||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||
pretrained_model_path, torch_dtype=torch.float16, use_safetensors=True)
|
||||
else:
|
||||
raise NotImplementedError(f"Invalid pretrained_model_version: {pretrained_model_version}")
|
||||
|
||||
pipe = pipe.to(device, dtype=torch.float16)
|
||||
print(f"Loaded new {pretrained_model_version} model from: {pretrained_model_path}")
|
||||
|
||||
# Load textual_inversion embeddings:
|
||||
load_ti_embeddings(pipe, lora_save_path)
|
||||
|
||||
# TODO: why does this give key errors???
|
||||
#pipe.load_lora_weights(lora_save_path, weight_name='pytorch_lora_weights.safetensors')
|
||||
#pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
|
||||
#pipe.fuse_lora(lora_scale=lora_scale)
|
||||
#return pipe
|
||||
|
||||
assert os.path.exists(lora_save_path), f"Invalid lora_save_path: {lora_save_path}"
|
||||
text_encoder_0_path = os.path.join(
|
||||
lora_save_path, "text_encoder_lora_0"
|
||||
)
|
||||
text_encoder_1_path = os.path.join(
|
||||
lora_save_path, "text_encoder_lora_1"
|
||||
)
|
||||
if os.path.exists(
|
||||
text_encoder_0_path
|
||||
):
|
||||
pipe.text_encoder = PeftModel.from_pretrained(pipe.text_encoder, text_encoder_0_path)
|
||||
print(f"loaded text_encoder LoRA from: {text_encoder_0_path}")
|
||||
|
||||
if os.path.exists(
|
||||
text_encoder_1_path
|
||||
):
|
||||
pipe.text_encoder_2 = PeftModel.from_pretrained(pipe.text_encoder_2, text_encoder_1_path)
|
||||
print(f"loaded text_encoder LoRA from: {text_encoder_1_path}")
|
||||
|
||||
if is_lora:
|
||||
pipe.unet = PeftModel.from_pretrained(model = pipe.unet, model_id = lora_save_path)
|
||||
else:
|
||||
pipe.unet = pipe.unet.from_pretrained(lora_save_path)
|
||||
print(f"Successfully loaded full checkpoint for inference!")
|
||||
|
||||
pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
|
||||
|
||||
return pipe
|
||||
@@ -0,0 +1,182 @@
|
||||
from typing import Union, List, Optional
|
||||
from datetime import datetime
|
||||
from pydantic import BaseModel
|
||||
import json, time, os
|
||||
from typing import Literal
|
||||
from trainer.utils.utils import pick_best_gpu_id
|
||||
|
||||
class ModelPaths:
|
||||
def __init__(self):
|
||||
self.paths = {
|
||||
"BLIP": "./cache",
|
||||
"CLIP": "./cache",
|
||||
"SR": "./cache",
|
||||
"SD": "./models",
|
||||
}
|
||||
|
||||
def get_path(self, key):
|
||||
return self.paths.get(key, None)
|
||||
|
||||
def set_path(self, key, path):
|
||||
if key in self.paths:
|
||||
self.paths[key] = path
|
||||
|
||||
model_paths = ModelPaths()
|
||||
|
||||
# Default download urls in case no local model is found:
|
||||
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-XL-v6/resolve/main/juggernautXL_version6Rundiffusion.safetensors"
|
||||
SDXL_URL = "https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_base_1.0_0.9vae.safetensors"
|
||||
SD15_URL = "https://huggingface.co/KamCastle/jugg/resolve/main/juggernaut_reborn.safetensors"
|
||||
|
||||
pretrained_models = {
|
||||
"sdxl": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SDXL_URL)), "url": SDXL_URL, "version": "sdxl"},
|
||||
"sd15": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SD15_URL)), "url": SD15_URL, "version": "sd15"}
|
||||
}
|
||||
|
||||
class TrainingConfig(BaseModel):
|
||||
lora_training_urls: str
|
||||
concept_mode: Literal["face", "style", "object"]
|
||||
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
|
||||
caption_model: Literal["gpt4-v", "blip"] = "blip"
|
||||
sd_model_version: Literal["sdxl", "sd15", None] = None
|
||||
ckpt_path: str = None # optional hardcoded checkpoint path
|
||||
pretrained_model: dict = None
|
||||
seed: Union[int, None] = None
|
||||
resolution: int = 512
|
||||
validation_img_size: Optional[Union[int, List[int]]] = None # [width, height], target_n_pixels ** 0.5 or None
|
||||
train_img_size: List[int] = None
|
||||
train_aspect_ratio: float = None
|
||||
train_batch_size: int = 4
|
||||
num_train_epochs: int = 10000
|
||||
max_train_steps: int = 360
|
||||
checkpointing_steps: int = 10000
|
||||
gradient_accumulation_steps: int = 1
|
||||
is_lora: bool = True
|
||||
|
||||
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
|
||||
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
|
||||
unet_lr: float = 1.0e-3
|
||||
prodigy_d_coef: float = 1.0
|
||||
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
|
||||
lora_weight_decay: float = 0.002
|
||||
|
||||
ti_lr: float = 1e-3
|
||||
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
|
||||
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
|
||||
ti_weight_decay: float = 0.0
|
||||
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
|
||||
freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done
|
||||
|
||||
cond_reg_w: float = 0.0e-5
|
||||
tok_cond_reg_w: float = 0.0e-5
|
||||
tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens
|
||||
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std
|
||||
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
|
||||
|
||||
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
|
||||
snr_gamma: float = 5.0
|
||||
lora_alpha_multiplier: float = 1.0
|
||||
lora_rank: int = 12
|
||||
use_dora: bool = False
|
||||
|
||||
left_right_flip_augmentation: bool = True
|
||||
augment_imgs_up_to_n: int = 40
|
||||
mask_target_prompts: Union[None, str] = None
|
||||
crop_based_on_salience: bool = True
|
||||
use_face_detection_instead: bool = False # use a different model (not CLIPSeg) to generate face masks
|
||||
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
|
||||
n_sample_imgs: int = 4
|
||||
name: str = None
|
||||
output_dir: str = "eden_lora_training_runs"
|
||||
debug: bool = False
|
||||
allow_tf32: bool = True
|
||||
disable_ti: bool = False
|
||||
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
|
||||
n_tokens: int = 2
|
||||
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
|
||||
token_dict: dict = {"TOK": "<s0><s1>"}
|
||||
device: str = "cuda:0"
|
||||
crops_coords_top_left_h: int = 0
|
||||
crops_coords_top_left_w: int = 0
|
||||
do_cache: bool = True
|
||||
unet_learning_rate: float = 1.0
|
||||
lr_num_cycles: int = 1
|
||||
lr_power: float = 1.0
|
||||
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
|
||||
dataloader_num_workers: int = 0
|
||||
training_attributes: dict = {}
|
||||
aspect_ratio_bucketing: bool = False
|
||||
start_time: float = 0.0
|
||||
job_time: float = 0.0
|
||||
"""
|
||||
For text encoder lora training, the trigger variable is: text_encoder_lora_optimizer
|
||||
if text_encoder_lora_optimizer is not None then everything else is used.
|
||||
Else the other variables are ignored.
|
||||
"""
|
||||
text_encoder_lora_optimizer: Union[None, Literal["adamw"]] = None
|
||||
text_encoder_lora_lr: float = 1.0e-5
|
||||
txt_encoders_lr_warmup_steps: int = 200
|
||||
text_encoder_lora_weight_decay: float = 1.0e-5
|
||||
text_encoder_lora_rank: int = 16
|
||||
|
||||
def __init__(self, **data):
|
||||
super().__init__(**data)
|
||||
|
||||
if not self.ckpt_path:
|
||||
self.pretrained_model = pretrained_models[self.sd_model_version]
|
||||
else:
|
||||
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
|
||||
|
||||
# add some metrics to the foldername:
|
||||
lora_str = "dora" if self.use_dora else "lora"
|
||||
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
|
||||
|
||||
if not self.name:
|
||||
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
|
||||
|
||||
self.output_dir = self.output_dir + f"/{self.name}/" + f"{timestamp_short}-{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
|
||||
os.makedirs(self.output_dir, exist_ok=True)
|
||||
|
||||
if self.seed is None:
|
||||
self.seed = int(time.time())
|
||||
|
||||
if self.unet_lr_warmup_steps is None:
|
||||
self.unet_lr_warmup_steps = self.max_train_steps
|
||||
|
||||
if self.concept_mode == "face":
|
||||
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
|
||||
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
|
||||
self.mask_target_prompts = "face"
|
||||
#self.use_face_detection_instead = True
|
||||
|
||||
if not self.sample_imgs_lora_scale:
|
||||
if self.sd_model_version == "sdxl":
|
||||
self.sample_imgs_lora_scale = 0.7
|
||||
else:
|
||||
self.sample_imgs_lora_scale = 0.85
|
||||
|
||||
if self.use_dora:
|
||||
print(f"Disabling L1 penalty and LoRA weight decay for DORA training.")
|
||||
self.l1_penalty = 0.0
|
||||
self.lora_weight_decay = 0.0
|
||||
self.text_encoder_lora_weight_decay = 0.0
|
||||
|
||||
# build the inserting_list_tokens and token dict using n_tokens:
|
||||
inserting_list_tokens = [f"<s{i}>" for i in range(self.n_tokens)]
|
||||
self.inserting_list_tokens = inserting_list_tokens
|
||||
self.token_dict = {"TOK": "".join(inserting_list_tokens)}
|
||||
|
||||
gpu_id = pick_best_gpu_id()
|
||||
self.device = f'cuda:{gpu_id}'
|
||||
self.start_time = time.time()
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, file_path: str):
|
||||
with open(file_path, 'r') as f:
|
||||
data = json.load(f)
|
||||
|
||||
return cls(**data)
|
||||
|
||||
def save_as_json(self, file_path: str) -> None:
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(self.dict(), f, indent=4)
|
||||
+168
-148
@@ -1,171 +1,191 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import PIL
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
from typing import Union
|
||||
from dataclasses import dataclass
|
||||
import torchvision.transforms as transforms
|
||||
import numpy as np
|
||||
import torch
|
||||
from typing import Tuple, Dict, List
|
||||
|
||||
@dataclass
|
||||
class ImageSize:
|
||||
width: int
|
||||
height: int
|
||||
def prepare_image(
|
||||
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
|
||||
) -> torch.Tensor:
|
||||
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
|
||||
image = pipe.image_processor.preprocess(pil_image)
|
||||
return image
|
||||
|
||||
default_tokenizer_kwargs = dict(
|
||||
padding="max_length",
|
||||
max_length=77,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt"
|
||||
)
|
||||
|
||||
default_image_transforms = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
|
||||
])
|
||||
|
||||
def convert_pil_mask_to_tensor(mask):
|
||||
arr = np.array(mask)
|
||||
def prepare_mask(
|
||||
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
|
||||
) -> torch.Tensor:
|
||||
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
|
||||
arr = np.array(pil_image.convert("L"))
|
||||
arr = arr.astype(np.float32) / 255.0
|
||||
return torch.tensor(arr).unsqueeze(0).unsqueeze(0)
|
||||
|
||||
class ImageCaptionDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
image_folder: str,
|
||||
csv_filename: str,
|
||||
mask_folder: Union[str, None] = None,
|
||||
validate_csv: bool = True,
|
||||
size: Union[ImageSize, None] = None,
|
||||
):
|
||||
super().__init__()
|
||||
assert os.path.exists(image_folder), f"Invalid image_folder: {image_folder}"
|
||||
assert os.path.exists(csv_filename), f"Invalid csv_filename: {csv_filename}"
|
||||
df = pd.read_csv(csv_filename)
|
||||
assert (
|
||||
"caption" in list(df.columns)
|
||||
), f"Expected column: 'caption' to exist but got columns: {list(df.columns)}"
|
||||
assert (
|
||||
"image_path" in list(df.columns)
|
||||
), f"Expected column: 'image_path' to exist but got columns: {list(df.columns)}"
|
||||
|
||||
self.csv_filename = csv_filename
|
||||
self.captions = df.caption.values
|
||||
self.image_path = df.image_path.values
|
||||
self.size = size
|
||||
self.image_folder = image_folder
|
||||
self.mask_folder = mask_folder
|
||||
|
||||
|
||||
if mask_folder is not None:
|
||||
assert (
|
||||
"mask_path" in list(df.columns)
|
||||
), f"Expected column: 'mask_path' to exist but got columns: {list(df.columns)}"
|
||||
self.mask_path = df.mask_path
|
||||
else:
|
||||
self.mask_path = None
|
||||
|
||||
if validate_csv:
|
||||
self.validate_csv()
|
||||
|
||||
def validate_csv(self):
|
||||
for idx in range(len(self.image_path)):
|
||||
filename = os.path.join(self.image_folder, self.image_path[idx])
|
||||
assert os.path.exists(
|
||||
filename
|
||||
), f"Invalid image path: {filename}\nPlease check your CSV file: {self.csv_filename}"
|
||||
|
||||
if self.mask_path is not None:
|
||||
for idx in range(len(self.mask_path)):
|
||||
filename = os.path.join(self.image_folder, self.image_path[idx])
|
||||
assert os.path.exists(
|
||||
filename
|
||||
), f"Invalid mask_path: {filename}\nPlease check your CSV file: {self.csv_filename}"
|
||||
|
||||
def __getitem__(self, idx: int) -> dict:
|
||||
filename = os.path.join(self.image_folder, self.image_path[idx])
|
||||
image = Image.open(filename)
|
||||
if self.size is not None:
|
||||
# inspired by dataset_and_utils.py -> prepare_image
|
||||
image = image.resize(
|
||||
(self.size.width, self.size.height),
|
||||
resample=Image.BICUBIC,
|
||||
reducing_gap=1,
|
||||
)
|
||||
image = image.convert("RGB")
|
||||
caption = self.captions[idx]
|
||||
|
||||
if self.mask_path is not None:
|
||||
image_width, image_height = image.size
|
||||
mask_filename = os.path.join(self.mask_folder, self.mask_path[idx])
|
||||
mask = Image.open(mask_filename)
|
||||
mask = mask.convert("L")
|
||||
mask = mask.resize(
|
||||
(image_width, image_height),
|
||||
resample=Image.BICUBIC,
|
||||
reducing_gap=1,
|
||||
)
|
||||
else:
|
||||
mask = None
|
||||
|
||||
return {"image": image, "caption": caption, "mask": mask}
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.captions)
|
||||
arr = np.expand_dims(arr, 0)
|
||||
image = torch.from_numpy(arr).unsqueeze(0)
|
||||
return image
|
||||
|
||||
|
||||
class PreprocessedDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
image_caption_dataset: ImageCaptionDataset,
|
||||
tokenizers: list,
|
||||
vae,
|
||||
text_encoders: Union[list, None] = None,
|
||||
cache: bool = True,
|
||||
tokenizer_kwargs: dict = default_tokenizer_kwargs,
|
||||
scale_vae_latents: bool = True
|
||||
data_dir: str,
|
||||
pipe,
|
||||
vae_encoder,
|
||||
do_cache: bool = False,
|
||||
size: List[int] = [512, 512],
|
||||
text_dropout: float = 0.0,
|
||||
aspect_ratio_bucketing: bool = False,
|
||||
train_batch_size: int = None, # required for aspect_ratio_bucketing
|
||||
substitute_caption_map: Dict[str, str] = {},
|
||||
):
|
||||
self.image_caption_dataset=image_caption_dataset
|
||||
self.tokenizers=tokenizers
|
||||
self.vae=vae
|
||||
self.text_encoders=text_encoders
|
||||
self.cache=cache
|
||||
self.tokenizer_kwargs=tokenizer_kwargs
|
||||
self.scale_vae_latents=scale_vae_latents
|
||||
|
||||
def __getitem__(self, idx: int):
|
||||
super().__init__()
|
||||
self.data_dir = data_dir
|
||||
self.csv_path = os.path.join(data_dir, "captions.csv")
|
||||
self.data = pd.read_csv(self.csv_path)
|
||||
|
||||
data = self.image_caption_dataset[idx]
|
||||
image, caption, mask = data["image"], data["caption"], data["mask"]
|
||||
self.captions = self.data["caption"]
|
||||
self.captions = self.captions.str.lower()
|
||||
for key, value in substitute_caption_map.items():
|
||||
self.captions = self.captions.str.replace(key.lower(), value)
|
||||
|
||||
tokenized_captions = []
|
||||
self.image_path = self.data["image_path"]
|
||||
|
||||
for tokenizer in self.tokenizers:
|
||||
tokenized_text = tokenizer(
|
||||
caption,
|
||||
**self.tokenizer_kwargs
|
||||
).input_ids.squeeze()
|
||||
tokenized_captions.append(
|
||||
tokenized_text
|
||||
if "mask_path" not in self.data.columns:
|
||||
self.mask_path = None
|
||||
else:
|
||||
self.mask_path = self.data["mask_path"]
|
||||
|
||||
self.pipe = pipe
|
||||
self.vae_encoder = vae_encoder
|
||||
self.vae_scaling_factor = self.vae_encoder.config.scaling_factor
|
||||
self.text_dropout = text_dropout
|
||||
self.size = size
|
||||
|
||||
if do_cache:
|
||||
print("Caching latents, masks and captions...\n")
|
||||
self.vae_latents = []
|
||||
self.masks = []
|
||||
self.do_cache = True
|
||||
|
||||
for idx in range(len(self.data)):
|
||||
if len(self.data) < 25:
|
||||
print(self.captions[idx])
|
||||
vae_latent, mask = self._process(idx)
|
||||
self.vae_latents.append(vae_latent)
|
||||
self.masks.append(mask)
|
||||
|
||||
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
|
||||
del self.vae_encoder
|
||||
else:
|
||||
self.do_cache = False
|
||||
|
||||
if aspect_ratio_bucketing:
|
||||
print("Using aspect ratio bucketing.")
|
||||
assert train_batch_size is not None, f"Please also provide a `train_batch_size` when you have set `aspect_ratio_bucketing == True`"
|
||||
from .utils.aspect_ratio_bucketing import BucketManager
|
||||
aspect_ratios = {}
|
||||
for idx in range(len(self.data)):
|
||||
aspect_ratios[idx] = Image.open(os.path.join(self.data_dir, self.image_path[idx])).size
|
||||
|
||||
self.bucket_manager = BucketManager(
|
||||
aspect_ratios = aspect_ratios,
|
||||
bsz = train_batch_size,
|
||||
debug=True
|
||||
)
|
||||
else:
|
||||
print("Not using aspect ratio bucketing.")
|
||||
self.bucket_manager = None
|
||||
|
||||
def get_aspect_ratio_bucketed_batch(self):
|
||||
assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()"
|
||||
indices, resolution = self.bucket_manager.get_batch()
|
||||
|
||||
print(f"Got bucket batch: {indices}, resolution: {resolution}")
|
||||
tok1, tok2, vae_latents, masks = [], [], [], []
|
||||
|
||||
for idx in indices:
|
||||
|
||||
if self.tokenizer_2 is None:
|
||||
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
|
||||
else:
|
||||
(t1, t2), v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
|
||||
tok2.append(t2.unsqueeze(0))
|
||||
|
||||
tok1.append(t1.unsqueeze(0))
|
||||
vae_latents.append(v.unsqueeze(0))
|
||||
masks.append(m.unsqueeze(0))
|
||||
|
||||
tok1 = torch.cat(tok1, dim = 0)
|
||||
if self.tokenizer_2 is None:
|
||||
pass
|
||||
else:
|
||||
tok2 = torch.cat(tok2, dim = 0)
|
||||
vae_latents = torch.cat(vae_latents, dim = 0)
|
||||
masks = torch.cat(masks, dim = 0)
|
||||
|
||||
if self.tokenizer_2 is None:
|
||||
return (tok1, None), vae_latents, masks
|
||||
else:
|
||||
return (tok1, tok2), vae_latents, masks
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.data)
|
||||
|
||||
@torch.no_grad()
|
||||
def _process(
|
||||
self, idx: int, bucketing_resolution: tuple = None
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
|
||||
image_path = self.image_path[idx]
|
||||
image_path = os.path.join(self.data_dir, image_path)
|
||||
image = PIL.Image.open(image_path).convert("RGB")
|
||||
if bucketing_resolution is None:
|
||||
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to(
|
||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
else:
|
||||
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
|
||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
image_tensor = default_image_transforms(image).unsqueeze(0).to(self.vae.device, dtype = self.vae.dtype)
|
||||
vae_latent = self.vae_encoder.encode(image).latent_dist
|
||||
dummy_vae_latent = vae_latent.sample()
|
||||
|
||||
if mask is not None:
|
||||
mask = convert_pil_mask_to_tensor(mask)
|
||||
if self.mask_path is None:
|
||||
mask = torch.ones_like(
|
||||
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
# raise ValueError(image_tensor.mean(), image_tensor.var())
|
||||
vae_latent = self.vae.encode(image_tensor).latent_dist.sample()
|
||||
if self.scale_vae_latents:
|
||||
vae_latent = vae_latent * self.vae.config.scaling_factor
|
||||
else:
|
||||
mask_path = self.mask_path[idx]
|
||||
mask_path = os.path.join(self.data_dir, mask_path)
|
||||
mask = PIL.Image.open(mask_path)
|
||||
mask = prepare_mask(mask, self.size[0], self.size[1]).to(
|
||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
||||
)
|
||||
|
||||
mask_dtype = mask.dtype
|
||||
mask = mask.float()
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask, size=(dummy_vae_latent.shape[-2], dummy_vae_latent.shape[-1]), mode="nearest"
|
||||
)
|
||||
mask = mask.to(dtype=mask_dtype)
|
||||
mask = mask.repeat(1, dummy_vae_latent.shape[1], 1, 1)
|
||||
|
||||
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
|
||||
|
||||
return vae_latent, mask.squeeze()
|
||||
|
||||
def __getitem__(
|
||||
self, idx: int, bucketing_resolution:tuple = None
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
|
||||
|
||||
if self.do_cache:
|
||||
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
|
||||
return self.captions[idx], vae_latent.squeeze(), self.masks[idx]
|
||||
else: # This code pathway has not been tested in a long time and might be broken
|
||||
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution)
|
||||
vae_latent = vae_latent.sample() * self.vae_scaling_factor
|
||||
return caption, vae_latent.squeeze(), mask
|
||||
|
||||
return {
|
||||
"tokenized_captions": tokenized_captions,
|
||||
"vae_latent": vae_latent.squeeze(),
|
||||
"mask": mask
|
||||
}
|
||||
|
||||
def __len__(self):
|
||||
return len(self.image_caption_dataset)
|
||||
|
||||
@@ -0,0 +1,497 @@
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
import PIL
|
||||
from tqdm import tqdm
|
||||
from typing import List, Optional, Dict
|
||||
from safetensors.torch import save_file, safe_open
|
||||
import matplotlib.pyplot as plt
|
||||
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
|
||||
|
||||
class TokenEmbeddingsHandler:
|
||||
def __init__(self, text_encoders, tokenizers):
|
||||
self.text_encoders = text_encoders
|
||||
self.tokenizers = tokenizers
|
||||
|
||||
self.train_ids: Optional[torch.Tensor] = None
|
||||
self.inserting_toks: Optional[List[str]] = None
|
||||
self.embeddings_settings = {}
|
||||
|
||||
self.target_prompt = ""
|
||||
self.token_regularizer = None
|
||||
|
||||
def make_embeddings_trainable(self):
|
||||
"""
|
||||
Sets requires_grad to True for specific indices directly in the embeddings weight tensor.
|
||||
"""
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
|
||||
# Directly accessing and modifying the original weights tensor
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
|
||||
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
|
||||
|
||||
def get_trainable_embeddings(self):
|
||||
return self.get_embeddings_and_tokens(self.train_ids)
|
||||
|
||||
def get_embeddings_and_tokens(self, indices):
|
||||
"""
|
||||
Get the embeddings and tokens for the given indices using PyTorch indexing.
|
||||
This version avoids detaching the original tensor and returns a view into the original
|
||||
weights tensor whenever possible.
|
||||
"""
|
||||
embeddings, tokens = {}, {}
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
|
||||
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
|
||||
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
|
||||
|
||||
# Directly access the embedding weights without detaching
|
||||
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
|
||||
embeddings[f'txt_encoder_{idx}'] = token_embeddings
|
||||
|
||||
# Get all corresponding tokens for these embeddings
|
||||
token_list = self.tokenizers[idx].convert_ids_to_tokens(indices)
|
||||
tokens[f'txt_encoder_{idx}'] = token_list
|
||||
|
||||
return embeddings, tokens
|
||||
|
||||
def visualize_random_token_embeddings(self, output_dir, n = 6, token_list = None):
|
||||
"""
|
||||
Visualize the embeddings of n random tokens from each text encoder
|
||||
"""
|
||||
if token_list is not None:
|
||||
# Convert tokens to indices using the first tokenizer
|
||||
indices = self.tokenizers[0].convert_tokens_to_ids(token_list)
|
||||
else:
|
||||
# Randomly select indices
|
||||
n_tokens = len(self.text_encoders[0].text_model.embeddings.token_embedding.weight.data)
|
||||
indices = np.random.randint(0, n_tokens, n)
|
||||
|
||||
embeddings, tokens = self.get_embeddings_and_tokens(indices)
|
||||
|
||||
# Visualize the embeddings:
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
for i in range(n):
|
||||
token = tokens[f'txt_encoder_{idx}'][i]
|
||||
# Strip any backslashes from the token name:
|
||||
token = token.replace("/", "_")
|
||||
embedding = embeddings[f'txt_encoder_{idx}'][i]
|
||||
plot_torch_hist(embedding, 0, os.path.join(output_dir, 'ti_embeddings') , f"frozen_enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'green')
|
||||
|
||||
def find_nearest_tokens(self, query_embedding, tokenizer, text_encoder, idx, distance_metric, top_k = 5):
|
||||
# given a query embedding, compute the distance to all embeddings in the text encoder
|
||||
# and return the top_k closest tokens
|
||||
|
||||
assert distance_metric in ["l2", "cosine"], "distance_metric should be either 'l2' or 'cosine'"
|
||||
|
||||
# get all non-optimized embeddings:
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[index_no_updates]
|
||||
|
||||
# compute the distance between the query embedding and all embeddings:
|
||||
if distance_metric == "l2":
|
||||
diff = (embeddings - query_embedding.unsqueeze(0))**2
|
||||
distances = diff.sum(-1)
|
||||
distances, indices = torch.topk(distances, top_k, dim=0, largest=False)
|
||||
elif distance_metric == "cosine":
|
||||
distances = F.cosine_similarity(embeddings, query_embedding.unsqueeze(0), dim=-1)
|
||||
distances, indices = torch.topk(distances, top_k, dim=0, largest=True)
|
||||
|
||||
nearest_tokens = tokenizer.convert_ids_to_tokens(indices)
|
||||
return nearest_tokens, distances
|
||||
|
||||
|
||||
def print_token_info(self, distance_metric = "cosine"):
|
||||
print(f"----------- Closest tokens (distance_metric = {distance_metric}) --------------")
|
||||
current_token_embeddings, current_tokens = self.get_trainable_embeddings()
|
||||
idx = 0
|
||||
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if text_encoder is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
query_embeddings = current_token_embeddings[f'txt_encoder_{idx}']
|
||||
query_tokens = current_tokens[f'txt_encoder_{idx}']
|
||||
|
||||
for token_id, query_embedding in enumerate(query_embeddings):
|
||||
nearest_tokens, distances = self.find_nearest_tokens(query_embedding, tokenizer, text_encoder, idx, distance_metric)
|
||||
|
||||
# print the results:
|
||||
print(f"txt-encoder {idx}, token {token_id}: {query_tokens[token_id]}:")
|
||||
for i, (token, dist) in enumerate(zip(nearest_tokens, distances)):
|
||||
print(f"---> {distance_metric} of {dist:.4f}: {token}")
|
||||
|
||||
idx += 1
|
||||
|
||||
def plot_token_embeddings(self, example_tokens, output_folder = ".", x_range = [-0.05, 0.05]):
|
||||
print(f"Plotting embeddings for tokens: {example_tokens}")
|
||||
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
token_ids = tokenizer.convert_tokens_to_ids(example_tokens)
|
||||
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_ids].clone()
|
||||
|
||||
# plot the embeddings histogram:
|
||||
for token_name, embedding in zip(example_tokens, embeddings):
|
||||
plot_torch_hist(embedding, 0, output_folder, f"tok_{token_name}_{idx}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
|
||||
|
||||
idx += 1
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return self.text_encoders[0].dtype
|
||||
|
||||
def initialize_new_tokens(self,
|
||||
inserting_toks: List[str],
|
||||
starting_toks: Optional[List[str]] = None,
|
||||
seed: int = 0,
|
||||
):
|
||||
assert isinstance(
|
||||
inserting_toks, list
|
||||
), "inserting_toks should be a list of strings."
|
||||
assert all(
|
||||
isinstance(tok, str) for tok in inserting_toks
|
||||
), "All elements in inserting_toks should be strings."
|
||||
|
||||
print(f"Initializing new tokens: {inserting_toks}")
|
||||
self.inserting_toks = inserting_toks
|
||||
|
||||
seed_everything(seed)
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
print(f"Inserting new tokens into tokenizer-{idx}:")
|
||||
print(self.inserting_toks)
|
||||
|
||||
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
|
||||
tokenizer.add_special_tokens(special_tokens_dict)
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
|
||||
|
||||
# construct the indices for all the non-trainable embeddings:
|
||||
all_indices = torch.linspace(0, len(tokenizer) - 1, len(tokenizer), dtype=torch.long)
|
||||
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
inu[self.train_ids] = False
|
||||
self.non_train_ids = all_indices[inu]
|
||||
|
||||
# random initialization of new tokens
|
||||
std_token_embedding = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
|
||||
)
|
||||
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
|
||||
|
||||
if starting_toks is not None:
|
||||
assert len(starting_toks) == len(self.inserting_toks), "starting_toks should have the same length as inserting_toks"
|
||||
self.starting_ids = tokenizer.convert_tokens_to_ids(starting_toks)
|
||||
print(f"Copying embeddings from starting tokens {starting_toks} to new tokens {self.inserting_toks}")
|
||||
print(f"Starting ids: {self.starting_ids}")
|
||||
# copy the embeddings of the starting tokens to the new tokens
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
|
||||
else:
|
||||
std_multiplier = 1.0
|
||||
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
|
||||
current_std = init_embeddings.std(dim=1).mean()
|
||||
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
|
||||
|
||||
self.embeddings_settings[
|
||||
f"original_embeddings_{idx}"
|
||||
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
|
||||
|
||||
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
inu[self.train_ids] = False
|
||||
self.embeddings_settings[f"index_no_updates_{idx}"] = inu
|
||||
|
||||
idx += 1
|
||||
|
||||
def plot_tokenid(self, token_id, suffix = '', output_folder = ".", x_range = [-0.05, 0.05]):
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_id].clone()
|
||||
plot_torch_hist(embeddings, 0, output_folder, f"tok_{token_id}_{idx}_{suffix}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
|
||||
idx += 1
|
||||
|
||||
def get_conditioning_signals(self, config, pipe, captions):
|
||||
conditioning_signals = pipe.encode_prompt(
|
||||
prompt=captions,
|
||||
device=pipe.unet.device,
|
||||
num_images_per_prompt=1,
|
||||
do_classifier_free_guidance=True,
|
||||
negative_prompt=None,
|
||||
clip_skip=None,
|
||||
)
|
||||
|
||||
try: # sd15
|
||||
prompt_embeds, negative_prompt_embeds = conditioning_signals
|
||||
pooled_prompt_embeds, add_time_ids = None, None
|
||||
|
||||
except: # sdxl
|
||||
(
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = conditioning_signals
|
||||
|
||||
# Create Spatial-dimensional conditions.
|
||||
# I dont understand why, but I get better results hardcoding the original_size values...
|
||||
# original_size = (config.resolution, config.resolution)
|
||||
original_size = (1024, 1024)
|
||||
target_size = (config.resolution, config.resolution)
|
||||
|
||||
crops_coords_top_left = (
|
||||
config.crops_coords_top_left_h,
|
||||
config.crops_coords_top_left_w,
|
||||
)
|
||||
|
||||
if pipe.text_encoder_2 is None:
|
||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||
else:
|
||||
text_encoder_projection_dim = pipe.text_encoder_2.config.projection_dim
|
||||
|
||||
add_time_ids = pipe._get_add_time_ids(
|
||||
original_size,
|
||||
crops_coords_top_left,
|
||||
target_size,
|
||||
dtype=prompt_embeds.dtype,
|
||||
text_encoder_projection_dim=text_encoder_projection_dim,
|
||||
)
|
||||
|
||||
add_time_ids = add_time_ids.to(config.device, dtype=prompt_embeds.dtype).repeat(
|
||||
prompt_embeds.shape[0], 1
|
||||
)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, add_time_ids
|
||||
|
||||
def encode_text(self, text, config, pipe):
|
||||
prompt_embeds, pooled_prompt_embeds, add_time_ids = self.get_conditioning_signals(config, pipe, [text])
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
def compute_target_prompt_loss(self, target_prompt, prompt_embeds, pooled_prompt_embeds, config, pipe):
|
||||
"""
|
||||
Compute a distance loss between the prompt embeddings and the target prompt embeddings
|
||||
"""
|
||||
|
||||
if target_prompt != self.target_prompt:
|
||||
self.target_prompt = target_prompt
|
||||
self.target_prompt_embeds, self.target_pooled_prompt_embeds = self.encode_text(self.target_prompt, config, pipe)
|
||||
# detach the target prompt embeddings (we don't need gradients here, these are just static targets)
|
||||
self.target_prompt_embeds = self.target_prompt_embeds.detach()
|
||||
try:
|
||||
self.target_pooled_prompt_embeds = self.target_pooled_prompt_embeds.detach()
|
||||
except:
|
||||
self.target_pooled_prompt_embeds = None
|
||||
|
||||
# compute the losses:
|
||||
|
||||
# Replicate target embeddings to match the batch size of prompt_embeds
|
||||
batch_size = prompt_embeds.size(0)
|
||||
target = self.target_prompt_embeds.expand(batch_size, -1, -1)
|
||||
embeds_l2_loss = F.mse_loss(prompt_embeds, target)
|
||||
embeds_cosine_loss = 1.0 - F.cosine_similarity(prompt_embeds, target, dim=-1).mean()
|
||||
|
||||
loss = embeds_l2_loss + embeds_cosine_loss
|
||||
|
||||
if pooled_prompt_embeds is not None:
|
||||
target = self.target_pooled_prompt_embeds.expand(batch_size, -1)
|
||||
pooled_embeds_l2_loss = F.mse_loss(pooled_prompt_embeds, target)
|
||||
pooled_embeds_cosine_loss = 1.0 - F.cosine_similarity(pooled_prompt_embeds, target, dim=-1).mean()
|
||||
loss += 0.25 * (pooled_embeds_l2_loss + pooled_embeds_cosine_loss)
|
||||
|
||||
return loss
|
||||
|
||||
def pre_optimize_token_embeddings(self, config, pipe):
|
||||
"""
|
||||
Warmup the token embeddings by optimizing them without using the image denoiser,
|
||||
but simply using CLIP-txt and CLIP-img similarity losses
|
||||
|
||||
TODO: add CLIP-img similarity loss into this mix
|
||||
--> This requires loading the img-encoder part for each of the txt-encoders and figuring out the correct projection layer
|
||||
"""
|
||||
target_prompt = config.training_attributes["gpt_description"]
|
||||
|
||||
if config.token_warmup_steps <= 0 or not target_prompt:
|
||||
print("Skipping token embedding warmup.")
|
||||
return
|
||||
|
||||
print(f'Warming up token embeddings with prompt: {target_prompt}...')
|
||||
|
||||
# Setup the token optimizer:
|
||||
ti_parameters = []
|
||||
for text_encoder in self.text_encoders:
|
||||
if text_encoder is not None:
|
||||
text_encoder.train()
|
||||
for name, param in text_encoder.named_parameters():
|
||||
if "token_embedding" in name:
|
||||
param.requires_grad = True
|
||||
ti_parameters.append(param)
|
||||
|
||||
params_to_optimize_ti = [{
|
||||
"params": ti_parameters,
|
||||
"lr": config.ti_lr,
|
||||
"weight_decay": config.ti_weight_decay,
|
||||
}]
|
||||
|
||||
optimizer_ti = torch.optim.AdamW(
|
||||
params_to_optimize_ti,
|
||||
weight_decay=config.ti_weight_decay,
|
||||
)
|
||||
|
||||
token_string = config.token_dict["TOK"]
|
||||
|
||||
# TODO: check if some light prompt template augmentation is useful here to make the optimization more robust
|
||||
prompt_template = [
|
||||
'{}',
|
||||
'{}',
|
||||
'{}',
|
||||
#'a {}',
|
||||
#'{} image',
|
||||
#'a picture of {}',
|
||||
]
|
||||
|
||||
losses = {'concept_description_loss': [], 'covariance_tok_reg_loss': [], 'token_std_loss': []}
|
||||
for step in tqdm(range(config.token_warmup_steps)):
|
||||
if step % 30 == 0 and config.debug and 0: # disalbe this for now
|
||||
for i, token_index in enumerate(self.train_ids):
|
||||
self.plot_tokenid(token_index, suffix = f'token_{i}_{step}', output_folder = f'{config.output_dir}/token_opt')
|
||||
|
||||
# pick a random prompt template and inject the token string:
|
||||
prompt_to_optimize = np.random.choice(prompt_template).format(token_string)
|
||||
prompt_embeds, pooled_prompt_embeds = self.encode_text(prompt_to_optimize, config, pipe)
|
||||
|
||||
# Compute the target_prompt distance loss:
|
||||
loss = 0.2 * self.compute_target_prompt_loss(target_prompt, prompt_embeds, pooled_prompt_embeds, config, pipe)
|
||||
losses['concept_description_loss'].append(loss.item())
|
||||
|
||||
# Compute token regularization loss:
|
||||
loss, losses, _ = self.token_regularizer.apply_regularization(loss, losses, None, prompt_embeds, std_loss_w = 0.5)
|
||||
|
||||
# Backward pass:
|
||||
retain_graph = step < (config.token_warmup_steps - 1) # Retain graph for all but the last step
|
||||
loss.backward(retain_graph=retain_graph)
|
||||
|
||||
# zero out the gradients of the non-trained text-encoder embeddings
|
||||
for embedding_tensor in ti_parameters:
|
||||
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
|
||||
|
||||
optimizer_ti.step()
|
||||
self.fix_embedding_std(config.off_ratio_power)
|
||||
optimizer_ti.zero_grad()
|
||||
|
||||
if config.debug:
|
||||
plot_loss(losses, save_path=f'{config.output_dir}/token_warmup_loss.png')
|
||||
|
||||
def save_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
|
||||
assert (
|
||||
self.train_ids is not None
|
||||
), "Initialize new tokens before saving embeddings."
|
||||
|
||||
# Create a set of indices for the non-train_ids:
|
||||
self.not_train_ids = torch.linspace(0, len(self.tokenizers[0]) - 1, len(self.tokenizers[0]), dtype=torch.long)
|
||||
tensors = {}
|
||||
for idx, text_encoder in enumerate(self.text_encoders):
|
||||
if text_encoder is None:
|
||||
continue
|
||||
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
|
||||
0
|
||||
] == len(self.tokenizers[0]), "Tokenizers should be the same."
|
||||
new_token_embeddings = (
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids
|
||||
]
|
||||
)
|
||||
tensors[txt_encoder_keys[idx]] = new_token_embeddings
|
||||
|
||||
save_file(tensors, file_path)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return self.text_encoders[0].device
|
||||
|
||||
def fix_embedding_std(self, off_ratio_power=0.1):
|
||||
if off_ratio_power == 0.0:
|
||||
return
|
||||
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
||||
if text_encoder is None:
|
||||
idx += 1
|
||||
continue
|
||||
|
||||
# Get the standard deviation target and current embeddings.
|
||||
target_std = self.embeddings_settings[f"std_token_embedding_{idx}"]
|
||||
embeddings, _ = self.get_trainable_embeddings()
|
||||
new_embeddings = embeddings[f'txt_encoder_{idx}']
|
||||
assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!"
|
||||
|
||||
new_stds = new_embeddings.std(dim=1)
|
||||
#off_ratios = target_std.float() / new_stds.float()
|
||||
off_ratios = target_std / new_stds
|
||||
|
||||
# Check if off_ratios are within an acceptable range.
|
||||
if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1):
|
||||
# Convert the pytorch tensor into a list of python floats:
|
||||
off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3)
|
||||
print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...")
|
||||
|
||||
# Adjust embeddings using the computed ratios.
|
||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
||||
index_updates = ~index_no_updates
|
||||
multiplier_values = off_ratios**off_ratio_power
|
||||
multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings)
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values
|
||||
|
||||
idx += 1
|
||||
|
||||
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
|
||||
# Assuming new tokens are of the format <s_i>
|
||||
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
|
||||
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
|
||||
tokenizer.add_special_tokens(special_tokens_dict)
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
|
||||
assert self.train_ids is not None, "New tokens could not be converted to IDs."
|
||||
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||
self.train_ids
|
||||
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
|
||||
|
||||
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
|
||||
if not os.path.exists(file_path):
|
||||
file_path = file_path.replace(".pti", ".safetensors")
|
||||
if not os.path.exists(file_path):
|
||||
raise FileNotFoundError(f"{file_path} does not exist.")
|
||||
|
||||
with safe_open(file_path, framework="pt", device=self.device.type) as f:
|
||||
for idx in range(len(self.text_encoders)):
|
||||
text_encoder = self.text_encoders[idx]
|
||||
tokenizer = self.tokenizers[idx]
|
||||
if text_encoder is None:
|
||||
continue
|
||||
try:
|
||||
loaded_embeddings = f.get_tensor(txt_encoder_keys[idx])
|
||||
except:
|
||||
loaded_embeddings = f.get_tensor(f"text_encoders_{idx}")
|
||||
self._load_embeddings(loaded_embeddings, tokenizer, text_encoder)
|
||||
@@ -0,0 +1,489 @@
|
||||
import torch
|
||||
import os
|
||||
import random
|
||||
import shutil
|
||||
import json
|
||||
import gc
|
||||
import re
|
||||
from diffusers import EulerDiscreteScheduler
|
||||
|
||||
from trainer.utils.val_prompts import val_prompts
|
||||
from trainer.utils.utils import fix_prompt, replace_in_string
|
||||
from trainer.models import load_models
|
||||
from .checkpoint import load_checkpoint, set_adapter_scales
|
||||
|
||||
from diffusers import (
|
||||
DDPMScheduler,
|
||||
EulerDiscreteScheduler,
|
||||
StableDiffusionPipeline,
|
||||
StableDiffusionXLPipeline,
|
||||
)
|
||||
|
||||
def load_model(pretrained_model: dict):
|
||||
if pretrained_model["version"] == "sd15":
|
||||
pipe = StableDiffusionPipeline.from_single_file(
|
||||
pretrained_model["path"], torch_dtype=torch.float16, use_safetensors=True
|
||||
)
|
||||
else:
|
||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||
pretrained_model["path"], torch_dtype=torch.float16, use_safetensors=True
|
||||
)
|
||||
|
||||
pipe = pipe.to("cuda", dtype=torch.float16)
|
||||
pipe.scheduler = EulerDiscreteScheduler.from_config(
|
||||
pipe.scheduler.config
|
||||
) # , timestep_spacing="trailing")
|
||||
|
||||
return pipe
|
||||
|
||||
|
||||
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
|
||||
"""
|
||||
This function is rather ugly, but implements a custom token-replacement policy we adopted at Eden:
|
||||
Basically you trigger the lora with a token "TOK" or "<concept>", and then this token gets replaced with the actual learned tokens
|
||||
"""
|
||||
|
||||
if "_no_token" in lora_path:
|
||||
return prompt
|
||||
|
||||
orig_prompt = prompt
|
||||
|
||||
# Helper function to read JSON
|
||||
def read_json_from_path(path):
|
||||
with open(path, "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
# Check existence of "special_params.json"
|
||||
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
|
||||
raise ValueError(
|
||||
"This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!"
|
||||
)
|
||||
|
||||
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
|
||||
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
|
||||
|
||||
trigger_text = training_args["training_attributes"]["trigger_text"]
|
||||
|
||||
try:
|
||||
lora_name = str(training_args["name"])
|
||||
except: # fallback for old loras that dont have the name field:
|
||||
lora_name = "concept"
|
||||
|
||||
lora_name_encapsulated = "<" + lora_name + ">"
|
||||
|
||||
try:
|
||||
mode = training_args["concept_mode"]
|
||||
except KeyError:
|
||||
try:
|
||||
mode = training_args["mode"]
|
||||
except KeyError:
|
||||
mode = "object"
|
||||
|
||||
# Handle different modes
|
||||
if mode != "style":
|
||||
replacements = {
|
||||
"<concept>": trigger_text,
|
||||
"<concepts>": trigger_text + "'s",
|
||||
lora_name_encapsulated: trigger_text,
|
||||
lora_name_encapsulated.lower(): trigger_text,
|
||||
lora_name: trigger_text,
|
||||
lora_name.lower(): trigger_text,
|
||||
}
|
||||
prompt = replace_in_string(prompt, replacements)
|
||||
if trigger_text not in prompt:
|
||||
prompt = trigger_text + ", " + prompt
|
||||
else:
|
||||
style_replacements = {
|
||||
"in the style of <concept>": "in the style of TOK",
|
||||
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
|
||||
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
|
||||
f"in the style of {lora_name}": "in the style of TOK",
|
||||
f"in the style of {lora_name.lower()}": "in the style of TOK",
|
||||
}
|
||||
prompt = replace_in_string(prompt, style_replacements)
|
||||
if "in the style of TOK" not in prompt:
|
||||
prompt = "in the style of TOK, " + prompt
|
||||
|
||||
# Final cleanup
|
||||
prompt = replace_in_string(
|
||||
prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"}
|
||||
)
|
||||
|
||||
if interpolation and mode != "style":
|
||||
prompt = "TOK, " + prompt
|
||||
|
||||
# Replace tokens based on token map
|
||||
prompt = replace_in_string(prompt, token_map)
|
||||
prompt = fix_prompt(prompt)
|
||||
|
||||
if verbose:
|
||||
print("-------------------------")
|
||||
print("Adjusted prompt for LORA:")
|
||||
print(orig_prompt)
|
||||
print("-- to:")
|
||||
print(prompt)
|
||||
print("-------------------------")
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
|
||||
def get_conditioning_signals(config, pipe, captions):
|
||||
conditioning_signals = pipe.encode_prompt(
|
||||
prompt=captions,
|
||||
device=pipe.unet.device,
|
||||
num_images_per_prompt=1,
|
||||
do_classifier_free_guidance=True,
|
||||
negative_prompt=None,
|
||||
clip_skip=None,
|
||||
)
|
||||
|
||||
try: # sd15
|
||||
prompt_embeds, negative_prompt_embeds = conditioning_signals
|
||||
pooled_prompt_embeds, add_time_ids = None, None
|
||||
|
||||
except: # sdxl
|
||||
(
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = conditioning_signals
|
||||
|
||||
# Create Spatial-dimensional conditions.
|
||||
# I dont understand why, but I get better results hardcoding the original_size values...
|
||||
# original_size = (config.resolution, config.resolution)
|
||||
original_size = (1024, 1024)
|
||||
target_size = (config.resolution, config.resolution)
|
||||
|
||||
crops_coords_top_left = (
|
||||
config.crops_coords_top_left_h,
|
||||
config.crops_coords_top_left_w,
|
||||
)
|
||||
|
||||
if pipe.text_encoder_2 is None:
|
||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||
else:
|
||||
text_encoder_projection_dim = pipe.text_encoder_2.config.projection_dim
|
||||
|
||||
add_time_ids = pipe._get_add_time_ids(
|
||||
original_size,
|
||||
crops_coords_top_left,
|
||||
target_size,
|
||||
dtype=prompt_embeds.dtype,
|
||||
text_encoder_projection_dim=text_encoder_projection_dim,
|
||||
)
|
||||
|
||||
add_time_ids = add_time_ids.to(config.device, dtype=prompt_embeds.dtype).repeat(
|
||||
prompt_embeds.shape[0], 1
|
||||
)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, add_time_ids
|
||||
|
||||
|
||||
def blend_conditions(
|
||||
embeds1,
|
||||
embeds2,
|
||||
lora_scale,
|
||||
token_scale_power=0.4, # adjusts the curve of the interpolation
|
||||
min_token_scale=0.5, # minimum token scale (corresponds to lora_scale = 0)
|
||||
token_scale=None,
|
||||
verbose=1,
|
||||
):
|
||||
"""
|
||||
using lora_scale, apply linear interpolation between two sets of embeddings
|
||||
"""
|
||||
try: # sdxl:
|
||||
c1, uc1, pc1, puc1 = embeds1
|
||||
c2, uc2, pc2, puc2 = embeds2
|
||||
except: # sd15:
|
||||
c1, uc1 = embeds1
|
||||
c2, uc2 = embeds2
|
||||
pc1, pc2, puc1, puc2 = None, None, None, None
|
||||
|
||||
if token_scale is None: # compute the token_scale based on lora_scale:
|
||||
token_scale = lora_scale**token_scale_power
|
||||
# rescale the [0,1] range to [min_token_scale, 1] range:
|
||||
token_scale = min_token_scale + (1 - min_token_scale) * token_scale
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
f"Setting token_scale to {token_scale:.2f} (lora_scale = {lora_scale:.2f}, power = {token_scale_power})"
|
||||
)
|
||||
|
||||
try:
|
||||
c = (1 - token_scale) * c1 + token_scale * c2
|
||||
uc = (1 - token_scale) * uc1 + token_scale * uc2
|
||||
try:
|
||||
pc = (1 - token_scale) * pc1 + token_scale * pc2
|
||||
puc = (1 - token_scale) * puc1 + token_scale * puc2
|
||||
except:
|
||||
pc, puc = None, None
|
||||
|
||||
embeds = (c, uc, pc, puc)
|
||||
except:
|
||||
print(
|
||||
f"Error in blending conditions for toking interpolation, falling back to embeds2"
|
||||
)
|
||||
token_scale = 1.0
|
||||
embeds = (c2, uc2, pc2, puc2)
|
||||
|
||||
return embeds, token_scale
|
||||
|
||||
|
||||
def encode_prompt_advanced(
|
||||
pipe,
|
||||
lora_path,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
lora_scale,
|
||||
guidance_scale,
|
||||
token_scale=None,
|
||||
concept_mode=None,
|
||||
):
|
||||
"""
|
||||
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
|
||||
This allows interpolating the strength of the trained token in the final image.
|
||||
"""
|
||||
if lora_path:
|
||||
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
|
||||
else:
|
||||
lora_prompt = prompt
|
||||
|
||||
if concept_mode == "face":
|
||||
replace_str = "person"
|
||||
elif concept_mode == "object":
|
||||
replace_str = "object"
|
||||
else:
|
||||
replace_str = ""
|
||||
|
||||
zero_prompt = prompt.replace("<concept>", replace_str)
|
||||
zero_prompt = fix_prompt(zero_prompt)
|
||||
|
||||
print(f"Embedding lora prompt: {lora_prompt}")
|
||||
print(f"Embedding zero prompt: {zero_prompt}")
|
||||
|
||||
try: # sdxl:
|
||||
embeds = pipe.encode_prompt(
|
||||
lora_prompt,
|
||||
do_classifier_free_guidance=guidance_scale > 1,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
|
||||
zero_embeds = pipe.encode_prompt(
|
||||
zero_prompt,
|
||||
do_classifier_free_guidance=guidance_scale > 1,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
|
||||
except: # sd15:
|
||||
embeds = pipe.encode_prompt(lora_prompt, pipe.device, 1, True, negative_prompt)
|
||||
|
||||
zero_embeds = pipe.encode_prompt(
|
||||
zero_prompt, pipe.device, 1, True, negative_prompt
|
||||
)
|
||||
|
||||
embeds, token_scale = blend_conditions(
|
||||
zero_embeds, embeds, lora_scale, token_scale=token_scale
|
||||
)
|
||||
|
||||
return embeds
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def render_images(
|
||||
render_size,
|
||||
lora_path,
|
||||
train_step,
|
||||
seed,
|
||||
is_lora,
|
||||
pretrained_model,
|
||||
lora_scale,
|
||||
n_steps=25,
|
||||
n_imgs=4,
|
||||
device="cuda:0",
|
||||
pipe = None,
|
||||
checkpoint_folder: str = None,
|
||||
):
|
||||
if checkpoint_folder is not None:
|
||||
assert pipe is None, f"Expected either one of checkpoint_folder or pipe to be None. But got: checkpoint_folder: {checkpoint_folder} and pipe is not None"
|
||||
|
||||
if pipe is not None:
|
||||
assert checkpoint_folder is None, f"Expected either one of checkpoint_folder or pipe to be None. But got pipe is NOT None checkpoint_folder is: {checkpoint_folder}"
|
||||
|
||||
random.seed(seed)
|
||||
|
||||
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
|
||||
training_args = json.load(f)
|
||||
concept_mode = training_args["concept_mode"]
|
||||
|
||||
if concept_mode == "style":
|
||||
validation_prompts_raw = random.sample(val_prompts["style"], n_imgs)
|
||||
validation_prompts_raw[0] = ""
|
||||
|
||||
elif concept_mode == "face":
|
||||
validation_prompts_raw = random.sample(val_prompts["face"], n_imgs)
|
||||
validation_prompts_raw[0] = "<concept>"
|
||||
else:
|
||||
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
|
||||
validation_prompts_raw[0] = "<concept>"
|
||||
|
||||
if (
|
||||
checkpoint_folder is not None
|
||||
): # reload the entire pipeline from disk and load in the lora module
|
||||
print(f"Reloading checkpoint from disk: {checkpoint_folder}")
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
pipe = load_checkpoint(
|
||||
pretrained_model_version=pretrained_model["version"],
|
||||
pretrained_model_path=pretrained_model["path"],
|
||||
lora_save_path=checkpoint_folder,
|
||||
is_lora=is_lora,
|
||||
device=device,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
|
||||
else:
|
||||
assert pipe is not None
|
||||
training_scheduler = pipe.scheduler
|
||||
print(f"Using existing model for inference")
|
||||
print(
|
||||
f"Re-using training pipeline for inference, just swapping the scheduler.."
|
||||
)
|
||||
pipe.vae = pipe.vae.to(device).to(pipe.unet.dtype)
|
||||
|
||||
pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
|
||||
pipe.scheduler = EulerDiscreteScheduler.from_config(
|
||||
pipe.scheduler.config, timestep_spacing="trailing"
|
||||
)
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
|
||||
pipeline_args = {
|
||||
"num_inference_steps": n_steps,
|
||||
"guidance_scale": 8,
|
||||
"width": render_size[0],
|
||||
"height": render_size[1],
|
||||
}
|
||||
|
||||
for i in range(n_imgs):
|
||||
print(f"Rendering validation img with prompt: {validation_prompts_raw[i]}")
|
||||
c, uc, pc, puc = encode_prompt_advanced(
|
||||
pipe,
|
||||
lora_path,
|
||||
validation_prompts_raw[i],
|
||||
negative_prompt,
|
||||
lora_scale,
|
||||
guidance_scale=8,
|
||||
concept_mode=concept_mode,
|
||||
)
|
||||
|
||||
pipeline_args["prompt_embeds"] = c
|
||||
pipeline_args["negative_prompt_embeds"] = uc
|
||||
if pretrained_model["version"] == "sdxl":
|
||||
pipeline_args["pooled_prompt_embeds"] = pc
|
||||
pipeline_args["negative_pooled_prompt_embeds"] = puc
|
||||
|
||||
image = pipe(**pipeline_args, generator=generator).images[0]
|
||||
image.save(
|
||||
os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"),
|
||||
format="JPEG",
|
||||
quality=95,
|
||||
)
|
||||
|
||||
if checkpoint_folder is None:
|
||||
pipe.scheduler = training_scheduler
|
||||
pipe.vae = pipe.vae.to("cpu")
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# reset the adapter scales to 1.0
|
||||
pipe = set_adapter_scales(pipe, lora_scale = 1.0)
|
||||
|
||||
return validation_prompts_raw
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def render_images_eval(
|
||||
concept_mode: str,
|
||||
output_folder: str,
|
||||
render_size: tuple,
|
||||
checkpoint_folder: str,
|
||||
seed: int,
|
||||
is_lora: bool,
|
||||
pretrained_model: dict,
|
||||
trigger_text: str,
|
||||
lora_scale=0.7,
|
||||
n_steps=25,
|
||||
n_imgs=4,
|
||||
device="cuda:0",
|
||||
verbose: bool = True,
|
||||
):
|
||||
random.seed(seed)
|
||||
assert os.path.exists(output_folder), f"Invalid folder: {output_folder}"
|
||||
|
||||
if concept_mode == "style":
|
||||
validation_prompts_raw = random.sample(val_prompts["style"], n_imgs)
|
||||
validation_prompts_raw[0] = ""
|
||||
elif concept_mode == "face":
|
||||
validation_prompts_raw = random.sample(val_prompts["face"], n_imgs)
|
||||
validation_prompts_raw[0] = "<concept>"
|
||||
else:
|
||||
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
|
||||
validation_prompts_raw[0] = "<concept>"
|
||||
|
||||
print(f"Reloading entire pipeline from disk for eval...")
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
pipe = load_checkpoint(
|
||||
pretrained_model_version=pretrained_model["version"],
|
||||
pretrained_model_path=pretrained_model["path"],
|
||||
lora_save_path=checkpoint_folder,
|
||||
is_lora=is_lora,
|
||||
device=device,
|
||||
)
|
||||
|
||||
pipe.scheduler = EulerDiscreteScheduler.from_config(
|
||||
pipe.scheduler.config, timestep_spacing="trailing"
|
||||
)
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
|
||||
pipeline_args = {
|
||||
"num_inference_steps": n_steps,
|
||||
"guidance_scale": 8,
|
||||
"height": render_size[0],
|
||||
"width": render_size[1],
|
||||
}
|
||||
|
||||
filenames = []
|
||||
for i in range(n_imgs):
|
||||
print(f"Rendering validation img with prompt: {validation_prompts_raw[i]}")
|
||||
c, uc, pc, puc = encode_prompt_advanced(
|
||||
pipe,
|
||||
checkpoint_folder,
|
||||
validation_prompts_raw[i],
|
||||
negative_prompt,
|
||||
lora_scale,
|
||||
guidance_scale=8,
|
||||
concept_mode=concept_mode,
|
||||
)
|
||||
|
||||
pipeline_args["prompt_embeds"] = c
|
||||
pipeline_args["negative_prompt_embeds"] = uc
|
||||
if pretrained_model["version"] == "sdxl":
|
||||
pipeline_args["pooled_prompt_embeds"] = pc
|
||||
pipeline_args["negative_pooled_prompt_embeds"] = puc
|
||||
|
||||
image = pipe(**pipeline_args, generator=generator).images[0]
|
||||
filename = os.path.join(output_folder, f"{i}.jpg")
|
||||
image.save(
|
||||
filename,
|
||||
format="JPEG",
|
||||
quality=95,
|
||||
)
|
||||
filenames.append(filename)
|
||||
|
||||
return filenames, validation_prompts_raw
|
||||
+375
@@ -0,0 +1,375 @@
|
||||
import os
|
||||
import time
|
||||
import matplotlib.pyplot as plt
|
||||
import torch
|
||||
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
|
||||
from trainer.inference import get_conditioning_signals
|
||||
|
||||
def compute_snr(noise_scheduler, timesteps):
|
||||
"""
|
||||
Computes SNR as per
|
||||
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
|
||||
"""
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = alphas_cumprod**0.5
|
||||
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
|
||||
|
||||
# Expand the tensors.
|
||||
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
|
||||
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
|
||||
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
|
||||
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
|
||||
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
|
||||
|
||||
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
|
||||
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
|
||||
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
|
||||
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
|
||||
|
||||
# Compute SNR.
|
||||
snr = (alpha / sigma) ** 2
|
||||
return snr
|
||||
|
||||
def compute_grad_norm(parameters, norm_type = 2.0, foreach = None, error_if_nonfinite = False):
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
first_device = grads[0].device
|
||||
grouped_grads = _group_tensors_by_device_and_dtype([[g.detach() for g in grads]])
|
||||
norms = []
|
||||
for ((device, _), ([grads], _)) in grouped_grads.items():
|
||||
if (foreach is None or foreach) and _has_foreach_support(grads, device=device):
|
||||
norms.extend(torch._foreach_norm(grads, norm_type))
|
||||
elif foreach:
|
||||
raise RuntimeError(f'foreach=True was passed, but can\'t use the foreach API on {device.type} tensors')
|
||||
else:
|
||||
norms.extend([torch.linalg.vector_norm(g, norm_type) for g in grads])
|
||||
|
||||
total_norm = torch.linalg.vector_norm(torch.stack([norm.to(first_device) for norm in norms]), norm_type)
|
||||
|
||||
return total_norm
|
||||
|
||||
def compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps):
|
||||
# Get the unet prediction target depending on the prediction type:
|
||||
if noise_scheduler.config.prediction_type == "epsilon":
|
||||
target = noise
|
||||
elif noise_scheduler.config.prediction_type == "v_prediction":
|
||||
print(f"Using velocity prediction!")
|
||||
target = noise_scheduler.get_velocity(noisy_latent, noise, timesteps)
|
||||
else:
|
||||
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
|
||||
|
||||
loss = (model_pred - target).pow(2) * mask
|
||||
|
||||
if config.snr_gamma is None or config.snr_gamma == 0.0:
|
||||
# modulate loss by the inverse of the mask's mean value
|
||||
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
|
||||
mean_mask_values = mean_mask_values / mean_mask_values.mean()
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
|
||||
loss = loss.mean()
|
||||
|
||||
else:
|
||||
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
|
||||
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
|
||||
# This is discussed in Section 4.2 of the same paper.
|
||||
snr = compute_snr(noise_scheduler, timesteps)
|
||||
base_weight = (
|
||||
torch.stack([snr, config.snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
|
||||
)
|
||||
if noise_scheduler.config.prediction_type == "v_prediction":
|
||||
# Velocity objective needs to be floored to an SNR weight of one.
|
||||
mse_loss_weights = base_weight + 1
|
||||
else:
|
||||
# Epsilon and sample both use the same loss weights.
|
||||
mse_loss_weights = base_weight
|
||||
|
||||
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
|
||||
|
||||
# modulate loss by the inverse of the mask's mean value
|
||||
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
|
||||
mean_mask_values = mean_mask_values / mean_mask_values.mean()
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
|
||||
loss = loss.mean()
|
||||
|
||||
return loss
|
||||
|
||||
class ConditioningRegularizer:
|
||||
"""
|
||||
Regularizes:
|
||||
- the norms of the prompt_conditioning vectors
|
||||
- the statistics of the token embeddings.
|
||||
"""
|
||||
|
||||
def __init__(self, config, embedding_handler):
|
||||
self.config = config
|
||||
self.embedding_handler = embedding_handler
|
||||
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
|
||||
self.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
|
||||
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
|
||||
|
||||
self.distribution_regularizers = {}
|
||||
idx = 0
|
||||
for tokenizer, text_encoder in zip(embedding_handler.tokenizers, embedding_handler.text_encoders):
|
||||
if tokenizer is None:
|
||||
idx += 1
|
||||
continue
|
||||
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
|
||||
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
|
||||
idx += 1
|
||||
|
||||
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.003, pipe=None):
|
||||
noise_sigma = 0.0
|
||||
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
|
||||
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
|
||||
|
||||
if self.config.cond_reg_w > 0.0:
|
||||
reg_loss, regularization_norm_value = self._compute_regularization_loss(prompt_embeds)
|
||||
loss += self.config.cond_reg_w * reg_loss
|
||||
if prompt_embeds_norms is not None:
|
||||
prompt_embeds_norms['main'].append(regularization_norm_value.item())
|
||||
|
||||
if self.config.tok_cond_reg_w > 0.0 and pipe is not None:
|
||||
reg_loss, regularization_norm_value = self._compute_tok_regularization_loss(pipe)
|
||||
loss += self.config.tok_cond_reg_w * reg_loss
|
||||
if prompt_embeds_norms is not None:
|
||||
prompt_embeds_norms['reg'].append(regularization_norm_value.item())
|
||||
|
||||
if self.config.tok_cov_reg_w > 0.0:
|
||||
tot_reg_losses = []
|
||||
for key, distribution_regularizer in self.distribution_regularizers.items():
|
||||
reg_loss = distribution_regularizer.compute_covariance_loss(self.embedding_handler.get_trainable_embeddings()[0][key])
|
||||
tot_reg_losses.append(reg_loss)
|
||||
|
||||
mean_reg_loss = torch.stack(tot_reg_losses).mean()
|
||||
loss += self.config.tok_cov_reg_w * mean_reg_loss
|
||||
losses['covariance_tok_reg_loss'].append(mean_reg_loss.item())
|
||||
|
||||
if std_loss_w > 0.0:
|
||||
tot_std_losses = []
|
||||
for key, distribution_regularizer in self.distribution_regularizers.items():
|
||||
std_loss = distribution_regularizer.compute_std_loss(self.embedding_handler.get_trainable_embeddings()[0][key])
|
||||
tot_std_losses.append(std_loss)
|
||||
|
||||
mean_std_loss = torch.stack(tot_std_losses).mean()
|
||||
loss += std_loss_w * mean_std_loss
|
||||
losses['token_std_loss'].append(mean_std_loss.item())
|
||||
|
||||
return loss, losses, prompt_embeds_norms
|
||||
|
||||
def _compute_regularization_loss(self, prompt_embeds):
|
||||
conditioning_norms = prompt_embeds.norm(dim=-1).mean(dim=0)
|
||||
regularization_norm_value = conditioning_norms[2:].mean()
|
||||
reg_loss = (regularization_norm_value - self.target_norm).pow(2)
|
||||
return reg_loss, regularization_norm_value
|
||||
|
||||
def _compute_tok_regularization_loss(self, pipe):
|
||||
reg_captions = [caption.replace("TOK", self.token_replacement) for caption in self.reg_captions]
|
||||
reg_prompt_embeds, reg_pooled_prompt_embeds, reg_add_time_ids = get_conditioning_signals(
|
||||
self.config, pipe, reg_captions
|
||||
)
|
||||
|
||||
reg_conditioning_norms = reg_prompt_embeds.norm(dim=-1).mean(dim=0)
|
||||
regularization_norm_value = reg_conditioning_norms[2:].mean()
|
||||
reg_loss = (regularization_norm_value - self.target_norm).pow(2)
|
||||
|
||||
return reg_loss, regularization_norm_value
|
||||
|
||||
|
||||
class DistributionLoss(torch.nn.Module):
|
||||
"""
|
||||
Class to simplify the calculation of the covariance loss between the trained token embeddings and the pretrained embeddings.
|
||||
"""
|
||||
def __init__(self, pretrained_embeddings, dtype=torch.float32, outdir = None):
|
||||
super(DistributionLoss, self).__init__()
|
||||
print(f"Initialized a new DistributionLoss with shape: {pretrained_embeddings.shape}")
|
||||
self.dtype = dtype
|
||||
self.target_cov = self._calculate_covariance(pretrained_embeddings)
|
||||
self.target_stds = pretrained_embeddings.std(-1)
|
||||
self.target_stds_mean = self.target_stds.mean()
|
||||
self.target_stds_var = self.target_stds.std()**2 / self.target_stds.mean()
|
||||
|
||||
if outdir:
|
||||
# Plot a histogram of the stds:
|
||||
plt.figure()
|
||||
plt.hist(self.target_stds.detach().float().cpu().numpy(), bins=100)
|
||||
plt.title(f"stds of tokens (shape = {pretrained_embeddings.shape[0]} x {pretrained_embeddings.shape[1]})")
|
||||
plt.xlim(0, 0.02)
|
||||
plt.savefig(os.path.join(outdir, f"stds_histogram_{int(time.time()*100)}.png"))
|
||||
|
||||
def _calculate_covariance(self, embeddings):
|
||||
embeddings = embeddings.to(self.dtype)
|
||||
mean = embeddings.mean(0)
|
||||
embeddings_adjusted = embeddings - mean
|
||||
covariance = torch.mm(embeddings_adjusted.T, embeddings_adjusted) / (embeddings.size(0) - 1)
|
||||
return covariance
|
||||
|
||||
def compute_covariance_loss(self, new_embeddings):
|
||||
input_dtype = new_embeddings.dtype
|
||||
cov_new = self._calculate_covariance(new_embeddings)
|
||||
# Normalizing by the product of the dimensions of the covariance matrix.
|
||||
num_features = new_embeddings.size(1) # Assuming embeddings are of shape [n_samples, n_features]
|
||||
scale_factor = num_features * num_features
|
||||
loss = torch.norm(self.target_cov - cov_new, p='fro') / scale_factor
|
||||
return loss.to(input_dtype)
|
||||
|
||||
def compute_std_loss(self, new_embeddings):
|
||||
if new_embeddings.size(1) == 1:
|
||||
new_embeddings = new_embeddings.unsqueeze(0)
|
||||
|
||||
deviation_loss = ((self.target_stds_mean - new_embeddings.std(-1))**2 / self.target_stds_var).mean()
|
||||
|
||||
return deviation_loss
|
||||
|
||||
|
||||
|
||||
#######################################################################
|
||||
#######################################################################
|
||||
|
||||
## Everything below here is experimental stuff not yet fully functional:
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from torch.distributions import MultivariateNormal, Normal
|
||||
from torch.distributions.distribution import Distribution
|
||||
|
||||
class GaussianKDE(Distribution):
|
||||
def __init__(self, X, bw = 0.1):
|
||||
"""
|
||||
X : tensor (n, d)
|
||||
`n` points with `d` dimensions to which KDE will be fit
|
||||
bw : numeric
|
||||
bandwidth for Gaussian kernel
|
||||
"""
|
||||
self.X = X
|
||||
self.bw = bw
|
||||
self.dims = X.shape[-1]
|
||||
self.n = X.shape[0]
|
||||
self.mvn = MultivariateNormal(loc=torch.zeros(self.dims),
|
||||
covariance_matrix=torch.eye(self.dims))
|
||||
|
||||
def sample(self, num_samples):
|
||||
idxs = (np.random.uniform(0, 1, num_samples) * self.n).astype(int)
|
||||
norm = Normal(loc=self.X[idxs], scale=self.bw)
|
||||
return norm.sample()
|
||||
|
||||
def score_samples(self, Y, X=None):
|
||||
"""Returns the kernel density estimates of each point in `Y`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
Y : tensor (m, d)
|
||||
`m` points with `d` dimensions for which the probability density will
|
||||
be calculated
|
||||
X : tensor (n, d), optional
|
||||
`n` points with `d` dimensions to which KDE will be fit. Provided to
|
||||
allow batch calculations in `log_prob`. By default, `X` is None and
|
||||
all points used to initialize KernelDensityEstimator are included.
|
||||
|
||||
|
||||
Returns
|
||||
-------
|
||||
log_probs : tensor (m)
|
||||
log probability densities for each of the queried points in `Y`
|
||||
"""
|
||||
if X == None:
|
||||
X = self.X
|
||||
log_probs = torch.log(
|
||||
(self.bw**(-self.dims) *
|
||||
torch.exp(self.mvn.log_prob(
|
||||
(X.unsqueeze(1) - Y) / self.bw))).sum(dim=0) / self.n)
|
||||
|
||||
return log_probs
|
||||
|
||||
def log_prob(self, Y):
|
||||
"""Returns the total log probability of one or more points, `Y`, using
|
||||
a Multivariate Normal kernel fit to `X` and scaled using `bw`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
Y : tensor (m, d)
|
||||
`m` points with `d` dimensions for which the probability density will
|
||||
be calculated
|
||||
|
||||
Returns
|
||||
-------
|
||||
log_prob : numeric
|
||||
total log probability density for the queried points, `Y`
|
||||
"""
|
||||
|
||||
X_chunks = self.X.split(1000)
|
||||
Y_chunks = Y.split(1000)
|
||||
|
||||
log_prob = 0
|
||||
|
||||
for x in X_chunks:
|
||||
for y in Y_chunks:
|
||||
log_prob += self.score_samples(y, x).sum(dim=0)
|
||||
|
||||
return log_prob
|
||||
|
||||
|
||||
class DifferentiableHistogram:
|
||||
"""
|
||||
TODO fix this function
|
||||
|
||||
"""
|
||||
def __init__(self, x, bins=64, min_range=None, max_range=None, bandwidth=0.02):
|
||||
self.bins = bins
|
||||
self.bandwidth = bandwidth * (x.max() - x.min())
|
||||
|
||||
if min_range is None or max_range is None:
|
||||
self.min_range = x.min()
|
||||
self.max_range = x.max()
|
||||
else:
|
||||
self.min_range = min_range
|
||||
self.max_range = max_range
|
||||
|
||||
# Create bins
|
||||
self.bin_edges = torch.linspace(self.min_range, self.max_range, bins + 1).to(x.device)
|
||||
self.bin_centers = (self.bin_edges[:-1] + self.bin_edges[1:]) / 2.0
|
||||
|
||||
# Compute histogram using Gaussian smoothing with bandwidth
|
||||
distances = (x.unsqueeze(1) - self.bin_centers.unsqueeze(0)) / self.bandwidth
|
||||
weights = torch.exp(-0.5 * (distances ** 2))
|
||||
histogram = weights.sum(dim=0)
|
||||
|
||||
# Normalize to form a PDF
|
||||
self.pdf = histogram / histogram.sum()
|
||||
|
||||
# Plot the histogram for validation
|
||||
plt.figure()
|
||||
plt.plot(self.bin_centers.float().cpu().numpy(), self.pdf.float().cpu().numpy())
|
||||
plt.title(f"PDF of token embeddings (shape = {x.shape})")
|
||||
plt.xlim(0, x.max().item()*1.1)
|
||||
plt.savefig(f"pdf_histogram_{int(time.time()*100)}.png")
|
||||
plt.close()
|
||||
|
||||
def __call__(self, y):
|
||||
"""
|
||||
Compute the negative log likelihood for a given sample y.
|
||||
Arguments:
|
||||
- y: Tensor of shape (m,) for which to compute the loss.
|
||||
Returns:
|
||||
- loss: Scalar representing the negative log likelihood of sample y.
|
||||
"""
|
||||
y_distances = (y.unsqueeze(1) - self.bin_centers.unsqueeze(0)) / self.bandwidth
|
||||
y_weights = torch.exp(-0.5 * (y_distances ** 2))
|
||||
likelihoods = (self.pdf * y_weights).sum(dim=1)
|
||||
|
||||
nll = -torch.log(likelihoods).mean()
|
||||
return nll
|
||||
|
||||
|
||||
"""
|
||||
|
||||
Learned Notes on the token embeddings:
|
||||
shape = 49410, 768
|
||||
Computed means of [1, 768] = [0, 0, 0, 0,...]
|
||||
Computed stds of [1, 768] = [0.0139, 0.0139, 0.0139, ...]
|
||||
|
||||
Computed means of [49410, 1] = [0, 0, 0, 0,...]
|
||||
Computed stds of [49410, 1] = [0.0151, 0.0154, 0.0141, ..., 0.0396, 0.0150, 0.0148]
|
||||
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
import os
|
||||
import time
|
||||
import subprocess
|
||||
import torch
|
||||
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
|
||||
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
|
||||
# check if the model is already downloaded:
|
||||
if not os.path.exists(pretrained_model['path']):
|
||||
download_weights(pretrained_model['url'], pretrained_model['path'])
|
||||
|
||||
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
|
||||
|
||||
try:
|
||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||
sd_model_version = "sdxl"
|
||||
except:
|
||||
pipe = StableDiffusionPipeline.from_single_file(
|
||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||
sd_model_version = "sd15"
|
||||
|
||||
print(f"Loaded {sd_model_version} model!")
|
||||
|
||||
pipe = pipe.to(device, dtype=weight_dtype)
|
||||
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
|
||||
|
||||
vae = pipe.vae
|
||||
unet = pipe.unet
|
||||
tokenizer_one = pipe.tokenizer
|
||||
text_encoder_one = pipe.text_encoder
|
||||
|
||||
vae.requires_grad_(False)
|
||||
if keep_vae_float32:
|
||||
vae.to(device, dtype=torch.float32)
|
||||
else:
|
||||
vae.to(device, dtype=weight_dtype)
|
||||
if weight_dtype != torch.float32:
|
||||
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but may not be ideal for training..?")
|
||||
|
||||
unet.to(device, dtype=weight_dtype)
|
||||
text_encoder_one.requires_grad_(False)
|
||||
text_encoder_one.to(device, dtype=weight_dtype)
|
||||
|
||||
tokenizer_two = text_encoder_two = None
|
||||
if sd_model_version == "sdxl":
|
||||
tokenizer_two = pipe.tokenizer_2
|
||||
text_encoder_two = pipe.text_encoder_2
|
||||
text_encoder_two.requires_grad_(False)
|
||||
text_encoder_two.to(device, dtype=weight_dtype)
|
||||
|
||||
return (
|
||||
pipe,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
noise_scheduler,
|
||||
text_encoder_one,
|
||||
text_encoder_two,
|
||||
vae,
|
||||
unet,
|
||||
), sd_model_version
|
||||
|
||||
def download_weights(url, dest):
|
||||
start = time.time()
|
||||
print("downloading url: ", url)
|
||||
print("downloading to: ", dest, '...')
|
||||
|
||||
# Make sure the destination directory exists
|
||||
dest_dir = os.path.dirname(dest)
|
||||
if not os.path.exists(dest_dir):
|
||||
os.makedirs(dest_dir)
|
||||
|
||||
try:
|
||||
subprocess.check_call(["wget", "-q", "-O", dest, url])
|
||||
except subprocess.CalledProcessError as e:
|
||||
print("Error occurred while downloading:")
|
||||
print("Exit status:", e.returncode)
|
||||
print("Output:", e.output)
|
||||
except Exception as e:
|
||||
print("An unexpected error occurred:", e)
|
||||
|
||||
print(f"Downloading {url} took {time.time() - start} seconds")
|
||||
|
||||
|
||||
def print_trainable_parameters(model, model_name = ''):
|
||||
trainable_params = 0
|
||||
all_param = 0
|
||||
for name, param in model.named_parameters():
|
||||
all_param += param.numel()
|
||||
if param.requires_grad and "token_embedding" not in name:
|
||||
trainable_params += param.numel()
|
||||
line_delimiter = "#" * 80
|
||||
print(line_delimiter)
|
||||
print(
|
||||
f"Trainable {model_name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
|
||||
)
|
||||
print(line_delimiter)
|
||||
@@ -0,0 +1,276 @@
|
||||
from peft import LoraConfig, get_peft_model
|
||||
import torch
|
||||
import prodigyopt
|
||||
from typing import Iterable
|
||||
|
||||
def get_unet_optimizer(
|
||||
prodigy_d_coef: float,
|
||||
prodigy_growth_factor: float,
|
||||
lora_weight_decay: float,
|
||||
use_dora: bool,
|
||||
unet_trainable_params: Iterable,
|
||||
optimizer_name="prodigy"
|
||||
):
|
||||
## unet_trainable_params can be unet.parameters() or a list of lora params
|
||||
|
||||
# These learning rates will get overwritten in main.py:
|
||||
if optimizer_name == "adamw":
|
||||
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
|
||||
elif optimizer_name == "AdamW8bit":
|
||||
import bitsandbytes as bnb
|
||||
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
|
||||
elif optimizer_name == "prodigy":
|
||||
# Note: the specific settings of Prodigy seem to matter A LOT
|
||||
optimizer_unet = prodigyopt.Prodigy(
|
||||
unet_trainable_params,
|
||||
d_coef = prodigy_d_coef,
|
||||
lr=1.0,
|
||||
decouple=True,
|
||||
use_bias_correction=True,
|
||||
safeguard_warmup=True,
|
||||
weight_decay=lora_weight_decay if not use_dora else 0.0,
|
||||
betas=(0.9, 0.99),
|
||||
growth_rate=prodigy_growth_factor # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Invalid optimizer_name for unet: {optimizer_name}")
|
||||
|
||||
print(f"Created {optimizer_name} optimizer for unet!")
|
||||
return optimizer_unet
|
||||
|
||||
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
|
||||
def is_belong_to_blocks(key, blocks):
|
||||
try:
|
||||
for g in blocks:
|
||||
if g in key:
|
||||
return True
|
||||
return False
|
||||
except Exception as e:
|
||||
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
|
||||
|
||||
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
|
||||
if use_blora:
|
||||
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
|
||||
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
|
||||
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
|
||||
try:
|
||||
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
|
||||
|
||||
attns = [
|
||||
attn_processor_name.rsplit(".", 1)[0]
|
||||
for attn_processor_name, _ in unet.attn_processors.items()
|
||||
if is_belong_to_blocks(attn_processor_name, blocks)
|
||||
]
|
||||
|
||||
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
|
||||
return target_modules
|
||||
except Exception as e:
|
||||
raise type(e)(
|
||||
f"failed to get_target_modules, due to: {e}. "
|
||||
f"Please check the modules specified in --lora_unet_blocks are correct"
|
||||
)
|
||||
|
||||
|
||||
def get_unet_lora_parameters(
|
||||
lora_rank,
|
||||
lora_alpha_multiplier: float,
|
||||
lora_weight_decay: float,
|
||||
use_dora: bool,
|
||||
unet,
|
||||
pipe,
|
||||
):
|
||||
|
||||
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
|
||||
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
|
||||
|
||||
unet_lora_config = LoraConfig(
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_rank * lora_alpha_multiplier,
|
||||
init_lora_weights="gaussian",
|
||||
target_modules=target_modules,
|
||||
use_dora=use_dora,
|
||||
)
|
||||
|
||||
#unet.add_adapter(unet_lora_config)
|
||||
unet = get_peft_model(unet, unet_lora_config)
|
||||
pipe.unet = unet
|
||||
|
||||
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
|
||||
unet_trainable_params = [
|
||||
{
|
||||
"params": unet_lora_parameters,
|
||||
"weight_decay": lora_weight_decay if not use_dora else 0.0,
|
||||
},
|
||||
]
|
||||
return unet, unet_trainable_params, unet_lora_parameters
|
||||
|
||||
def get_textual_inversion_optimizer(
|
||||
text_encoders: list,
|
||||
textual_inversion_lr: float,
|
||||
textual_inversion_weight_decay,
|
||||
optimizer_name: str
|
||||
):
|
||||
text_encoder_parameters = []
|
||||
for text_encoder in text_encoders:
|
||||
if text_encoder is not None:
|
||||
text_encoder.train()
|
||||
for name, param in text_encoder.named_parameters():
|
||||
if "token_embedding" in name:
|
||||
#param.data = param.to(dtype=torch.float32)
|
||||
param.requires_grad = True
|
||||
text_encoder_parameters.append(param)
|
||||
print(f"Added {name} with shape {param.shape} to the trainable parameters")
|
||||
else:
|
||||
pass
|
||||
|
||||
params_to_optimize_ti = [
|
||||
{
|
||||
"params": text_encoder_parameters,
|
||||
"lr": textual_inversion_lr if (optimizer_name != "prodigy") else 1.0,
|
||||
"weight_decay":textual_inversion_weight_decay,
|
||||
},
|
||||
]
|
||||
|
||||
if optimizer_name == "prodigy":
|
||||
optimizer_ti = prodigyopt.Prodigy(
|
||||
params_to_optimize_ti,
|
||||
d_coef = 1.0,
|
||||
lr=1.0,
|
||||
decouple=True,
|
||||
use_bias_correction=True,
|
||||
safeguard_warmup=True,
|
||||
weight_decay=textual_inversion_weight_decay,
|
||||
betas=(0.9, 0.99),
|
||||
#growth_rate=1.5, # this slows down the lr_rampup
|
||||
)
|
||||
elif optimizer_name == "adamw":
|
||||
optimizer_ti = torch.optim.AdamW(
|
||||
params_to_optimize_ti,
|
||||
weight_decay=textual_inversion_weight_decay,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Invalid optimizer_name: '{optimizer_name}'")
|
||||
|
||||
print(f"Created {optimizer_name} optimizer for textual inversion!")
|
||||
return optimizer_ti, text_encoder_parameters
|
||||
|
||||
def get_text_encoder_lora_parameters(text_encoder, lora_rank, lora_alpha_multiplier, use_dora: bool):
|
||||
text_encoder_lora_config = LoraConfig(
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_rank * lora_alpha_multiplier,
|
||||
init_lora_weights="gaussian",
|
||||
target_modules=["k_proj", "q_proj", "v_proj", "out_proj"],
|
||||
use_dora=use_dora,
|
||||
)
|
||||
text_encoder_peft_model = get_peft_model(text_encoder, text_encoder_lora_config)
|
||||
text_encoder_lora_params = list(filter(lambda p: p.requires_grad, text_encoder_peft_model.parameters()))
|
||||
return text_encoder_peft_model, text_encoder_lora_params
|
||||
|
||||
def get_optimizer_and_peft_models_text_encoder_lora(
|
||||
text_encoders: list,
|
||||
lora_rank: int,
|
||||
lora_alpha_multiplier: float,
|
||||
use_dora: bool,
|
||||
optimizer_name: str,
|
||||
lora_lr: float,
|
||||
weight_decay: float
|
||||
):
|
||||
text_encoder_lora_parameters = []
|
||||
text_encoder_peft_models = []
|
||||
for text_encoder in text_encoders:
|
||||
if text_encoder is not None:
|
||||
text_encoder_peft_model, text_encoder_lora_params = get_text_encoder_lora_parameters(
|
||||
text_encoder=text_encoder,
|
||||
lora_rank=lora_rank,
|
||||
lora_alpha_multiplier=lora_alpha_multiplier,
|
||||
use_dora=use_dora
|
||||
)
|
||||
text_encoder_lora_parameters.extend(text_encoder_lora_params)
|
||||
text_encoder_peft_models.append(text_encoder_peft_model)
|
||||
else:
|
||||
text_encoder_peft_models.append(None)
|
||||
|
||||
if optimizer_name == "adamw":
|
||||
optimizer_text_encoder_lora = torch.optim.AdamW(
|
||||
text_encoder_lora_parameters,
|
||||
lr = lora_lr,
|
||||
weight_decay=weight_decay if not use_dora else 0.0
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Text encoder LoRA finetuning is not yet implemented for optimizer: {optimizer_name}")
|
||||
|
||||
return optimizer_text_encoder_lora, text_encoder_peft_models
|
||||
|
||||
|
||||
|
||||
def get_current_lr(optimizer):
|
||||
"""
|
||||
Helper class to get the current lr for various types of optimizers
|
||||
"""
|
||||
try:
|
||||
# Calculate the weighted average effective learning rate
|
||||
total_lr = 0
|
||||
total_params = 0
|
||||
for group in optimizer.param_groups:
|
||||
d = group['d']
|
||||
lr = group['lr']
|
||||
bias_correction = 1 # Default value
|
||||
if group['use_bias_correction']:
|
||||
beta1, beta2 = group['betas']
|
||||
k = group['k']
|
||||
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
|
||||
|
||||
effective_lr = d * lr * bias_correction
|
||||
|
||||
# Count the number of parameters in this group
|
||||
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
|
||||
total_lr += effective_lr * num_params
|
||||
total_params += num_params
|
||||
|
||||
if total_params == 0:
|
||||
return 0.0
|
||||
else: return total_lr / total_params
|
||||
except:
|
||||
return optimizer.param_groups[0]['lr']
|
||||
|
||||
|
||||
class OptimizerCollection:
|
||||
def __init__(
|
||||
self,
|
||||
optimizer_textual_inversion = None,
|
||||
optimizer_text_encoders = None,
|
||||
optimizer_unet = None,
|
||||
debug = False,
|
||||
):
|
||||
"""
|
||||
run operations on all the relevant optimizers with a single function call
|
||||
"""
|
||||
self.debug = debug
|
||||
self.optimizers = {
|
||||
'textual_inversion': optimizer_textual_inversion,
|
||||
'text_encoders': optimizer_text_encoders,
|
||||
'unet': optimizer_unet
|
||||
}
|
||||
|
||||
self.learning_rate_tracker = {'textual_inversion':[], 'text_encoders':[], 'unet':[]}
|
||||
|
||||
print("--> Initialized optimizers for:")
|
||||
for key in self.optimizers.keys():
|
||||
if self.optimizers[key] is not None:
|
||||
print(key)
|
||||
|
||||
def get_lr(self, key):
|
||||
return get_current_lr(self.optimizers[key])
|
||||
|
||||
def zero_grad(self):
|
||||
for key in self.optimizers.keys():
|
||||
if self.optimizers[key] is not None:
|
||||
self.optimizers[key].zero_grad()
|
||||
|
||||
def step(self):
|
||||
for key in self.optimizers.keys():
|
||||
if self.optimizers[key] is not None:
|
||||
self.optimizers[key].step()
|
||||
if self.debug:
|
||||
self.learning_rate_tracker[key].append(get_current_lr(self.optimizers[key]))
|
||||
|
||||
@@ -1,7 +1,3 @@
|
||||
# Have SwinIR upsample
|
||||
# Have BLIP auto caption
|
||||
# Have CLIPSeg auto mask concept
|
||||
|
||||
import gc
|
||||
import fnmatch
|
||||
import mimetypes
|
||||
@@ -25,6 +21,7 @@ import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from transformers import (
|
||||
BlipForConditionalGeneration,
|
||||
Blip2ForConditionalGeneration,
|
||||
@@ -36,14 +33,15 @@ from transformers import (
|
||||
Swin2SRImageProcessor,
|
||||
)
|
||||
|
||||
from io_utils import download_and_prep_training_data
|
||||
from trainer.utils.io import download_and_prep_training_data
|
||||
from trainer.utils.utils import fix_prompt
|
||||
from trainer.config import model_paths
|
||||
|
||||
import re
|
||||
import openai
|
||||
from openai import OpenAI
|
||||
from dotenv import load_dotenv
|
||||
load_dotenv()
|
||||
|
||||
try:
|
||||
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
||||
client = OpenAI(api_key=OPENAI_API_KEY)
|
||||
@@ -53,30 +51,20 @@ except:
|
||||
client = None
|
||||
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
|
||||
|
||||
MODEL_PATH = "./cache"
|
||||
MAX_GPT_PROMPTS = 40
|
||||
|
||||
import re
|
||||
def fix_prompt(prompt: str):
|
||||
# Remove extra commas and spaces, and fix space before punctuation
|
||||
prompt = re.sub(r"\s+", " ", prompt) # Replace multiple spaces with a single space
|
||||
prompt = re.sub(r",,", ",", prompt) # Replace double commas with a single comma
|
||||
prompt = re.sub(r"\s?,\s?", ", ", prompt) # Fix spaces around commas
|
||||
prompt = re.sub(r"\s?\.\s?", ". ", prompt) # Fix spaces around periods
|
||||
return prompt.strip() # Remove leading and trailing whitespace
|
||||
|
||||
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...)
|
||||
MIN_GPT_PROMPTS = 3
|
||||
MAX_GPT_PROMPTS = 50
|
||||
|
||||
def _find_files(pattern, dir="."):
|
||||
"""Return list of files matching pattern in a given directory, in absolute format.
|
||||
Unlike glob, this is case-insensitive.
|
||||
"""
|
||||
|
||||
rule = re.compile(fnmatch.translate(pattern), re.IGNORECASE)
|
||||
return [os.path.join(dir, f) for f in os.listdir(dir) if rule.match(f)]
|
||||
|
||||
|
||||
|
||||
def preprocess(
|
||||
config,
|
||||
working_directory,
|
||||
concept_mode,
|
||||
input_zip_path: Path,
|
||||
@@ -85,7 +73,6 @@ def preprocess(
|
||||
target_size: int,
|
||||
crop_based_on_salience: bool,
|
||||
use_face_detection_instead: bool,
|
||||
temp: float,
|
||||
left_right_flip_augmentation: bool = False,
|
||||
augment_imgs_up_to_n: int = 0,
|
||||
caption_model: str = "blip",
|
||||
@@ -93,7 +80,6 @@ def preprocess(
|
||||
) -> Path:
|
||||
|
||||
if os.path.exists(working_directory):
|
||||
print(f"working_directory {working_directory} already existed.. deleting and recreating!")
|
||||
shutil.rmtree(working_directory)
|
||||
os.makedirs(working_directory)
|
||||
|
||||
@@ -108,7 +94,8 @@ def preprocess(
|
||||
|
||||
download_and_prep_training_data(input_zip_path, TEMP_IN_DIR)
|
||||
|
||||
n_training_imgs, trigger_text, segmentation_prompt, captions = load_and_save_masks_and_captions(
|
||||
config = load_and_save_masks_and_captions(
|
||||
config,
|
||||
concept_mode,
|
||||
files=TEMP_IN_DIR,
|
||||
output_dir=TEMP_OUT_DIR,
|
||||
@@ -118,13 +105,12 @@ def preprocess(
|
||||
target_size=target_size,
|
||||
crop_based_on_salience=crop_based_on_salience,
|
||||
use_face_detection_instead=use_face_detection_instead,
|
||||
temp=temp,
|
||||
add_lr_flips = left_right_flip_augmentation,
|
||||
augment_imgs_up_to_n = augment_imgs_up_to_n,
|
||||
caption_model = caption_model
|
||||
)
|
||||
|
||||
return Path(TEMP_OUT_DIR), n_training_imgs, trigger_text, segmentation_prompt, captions
|
||||
return config, Path(TEMP_OUT_DIR)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -148,7 +134,7 @@ def swin_ir_sr(
|
||||
"""
|
||||
|
||||
model = Swin2SRForImageSuperResolution.from_pretrained(
|
||||
model_id, cache_dir=MODEL_PATH
|
||||
model_id, cache_dir = model_paths.get_path("SR")
|
||||
).to(device)
|
||||
processor = Swin2SRImageProcessor()
|
||||
|
||||
@@ -196,14 +182,15 @@ def clipseg_mask_generator(
|
||||
|
||||
if isinstance(target_prompts, str):
|
||||
print(
|
||||
f'Warning: only one target prompt "{target_prompts}" was given, so it will be used for all images'
|
||||
f'Using "{target_prompts}" as CLIP-segmentation prompt for all images.'
|
||||
)
|
||||
target_prompts = [target_prompts] * len(images)
|
||||
|
||||
model = None
|
||||
if any(target_prompts):
|
||||
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
||||
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
|
||||
model = CLIPSegForImageSegmentation.from_pretrained(
|
||||
model_id, cache_dir=MODEL_PATH
|
||||
model_id, cache_dir = model_paths.get_path("CLIP")
|
||||
).to(device)
|
||||
|
||||
masks = []
|
||||
@@ -237,79 +224,68 @@ def clipseg_mask_generator(
|
||||
|
||||
masks.append(mask)
|
||||
|
||||
# cleanup
|
||||
del model
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return masks
|
||||
|
||||
|
||||
|
||||
import textwrap
|
||||
def cleanup_prompts_with_chatgpt(
|
||||
prompts,
|
||||
concept_mode, # face / object / style
|
||||
seed, # seed for chatgpt reproducibility
|
||||
verbose = True):
|
||||
|
||||
if concept_mode == "object_injection":
|
||||
chat_gpt_prompt_1 = """
|
||||
I have a set of images, each containing the same concept / figure. I have the following (poor) descriptions for each image:
|
||||
"""
|
||||
|
||||
chat_gpt_prompt_2 = """
|
||||
I want you to:
|
||||
1. Find a good, short name/description of the single central concept that's in all the images. This [Concept Name] might eg already be present in the descriptions above, pick the most obvious name or words that would fit in all descriptions.
|
||||
2. Insert the text "TOK, [Concept Name]" into all the descriptions above by rephrasing them where needed to naturally contain the text TOK, [Concept Name] while keeping as much of the description as possible.
|
||||
|
||||
Reply by first stating the "Concept Name:", followed by an enumerated list (using "-") of all the revised "Descriptions:".
|
||||
"""
|
||||
|
||||
if concept_mode == "object":
|
||||
chat_gpt_prompt_1 = """
|
||||
Analyze a set of (poor) image descriptions each featuring the same concept, figure or thing.
|
||||
Tasks:
|
||||
1. Deduce a concise, fitting name for the concept that is visually descriptive (Concept Name).
|
||||
2. Substitute the concept in each description with the placeholder "TOK", rearranging or adjusting the text where needed. Hallucinate TOK into the description if necessary (but dont mention when doing so, simply provide the final description)!
|
||||
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of the placeholder string "TOK".
|
||||
The descriptions are:
|
||||
"""
|
||||
chat_gpt_prompt_1 = textwrap.dedent("""
|
||||
Analyze a set of (poor) image descriptions each featuring the same concept, figure or thing.
|
||||
Tasks:
|
||||
1. Deduce a concise (max 10 words) visual description of just the concept TOK (Concept Description), try to be as visually descriptive of TOK as possible!
|
||||
2. Substitute the concept in each description with the placeholder "TOK", rearranging or adjusting the text where needed. Hallucinate TOK into the description if necessary (but dont mention when doing so, simply provide the final description)!
|
||||
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of the placeholder string "TOK".
|
||||
The descriptions are:""")
|
||||
|
||||
chat_gpt_prompt_2 = """
|
||||
Respond with the chosen "Concept Name:" followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
|
||||
"""
|
||||
chat_gpt_prompt_2 = textwrap.dedent("""
|
||||
Respond with "Concept Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
|
||||
""")
|
||||
|
||||
elif concept_mode == "face":
|
||||
chat_gpt_prompt_1 = """
|
||||
Analyze a set of (poor) image descriptions, each featuring a person named TOK.
|
||||
Tasks:
|
||||
1. Rewrite each description, ensuring it refers only to a single person or character.
|
||||
2. Integrate "a photo of TOK" naturally into each description, rearranging or adjusting where needed.
|
||||
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
|
||||
The descriptions are:
|
||||
"""
|
||||
chat_gpt_prompt_1 = textwrap.dedent("""
|
||||
Analyze a set of (poor) image descriptions, each featuring a person named TOK.
|
||||
Tasks:
|
||||
1. Deduce a concise (max 10 words) visual description of TOK (TOK Description), try to be as visually descriptive of TOK as possible, always mention their skin color, hallucinate a basic description if necessary (eg black man with long beard).
|
||||
2. Rewrite each description, injecting "TOK" naturally into each description, adjusting where needed.
|
||||
3. Streamline each description to focus on the context and surroundings of TOK instead of the visual appearance of TOK's face. Ensure mandatory inclusion of "TOK".
|
||||
The descriptions are:""")
|
||||
|
||||
chat_gpt_prompt_2 = """
|
||||
Respond with "Concept Name: TOK" followed by a list (using "-") of all the revised descriptions, each mentioning "a photo of TOK".
|
||||
"""
|
||||
chat_gpt_prompt_2 = textwrap.dedent("""
|
||||
Respond with "TOK Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
|
||||
""")
|
||||
|
||||
elif concept_mode == "style":
|
||||
chat_gpt_prompt_1 = """
|
||||
Analyze a set of (poor) image descriptions, each featuring the same style named TOK.
|
||||
Tasks:
|
||||
1. Rewrite each description to focus solely on the TOK style.
|
||||
2. Integrate "in the style of TOK" naturally into each description, typically at the beginning.
|
||||
3. Summarize each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
|
||||
The descriptions are:
|
||||
"""
|
||||
chat_gpt_prompt_1 = textwrap.dedent("""
|
||||
Analyze a set of (poor) image descriptions, each featuring an example of a common aesthetic style named TOK.
|
||||
Tasks:
|
||||
1. Deduce a concise (max 7 words) visual description of the aesthetic style (Style Description).
|
||||
2. Rewrite each description to focus solely on the non-stylistic contents of the image like characters, objects, colors, scene, context etc but not the stylistic elements captured by TOK.
|
||||
3. Integrate "in the style of TOK" naturally into each description, typically at the beginning while summarizing each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
|
||||
The descriptions are:""")
|
||||
|
||||
chat_gpt_prompt_2 = """
|
||||
Respond with "Style Name: TOK" followed by a list (using "-") of all the revised descriptions, each mentioning "in the style of TOK".
|
||||
"""
|
||||
chat_gpt_prompt_2 = textwrap.dedent("""
|
||||
Respond with "Style Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "in the style of TOK".
|
||||
""")
|
||||
|
||||
final_chatgpt_prompt = chat_gpt_prompt_1 + "\n- " + "\n- ".join(prompts) + "\n\n" + chat_gpt_prompt_2
|
||||
final_chatgpt_prompt = chat_gpt_prompt_1 + "\n- " + "\n- ".join(prompts) + "\n" + chat_gpt_prompt_2
|
||||
print("Final chatgpt prompt:")
|
||||
print(final_chatgpt_prompt)
|
||||
print("--------------------------")
|
||||
print(f"Calling chatgpt with seed {seed}...")
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="gpt-4-1106-preview",
|
||||
model="gpt-4o",
|
||||
seed=seed,
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
@@ -326,67 +302,70 @@ def cleanup_prompts_with_chatgpt(
|
||||
# extract the final rephrased prompts from the response:
|
||||
prompts = []
|
||||
for line in gpt_completion.split("\n"):
|
||||
if line.startswith("-"):
|
||||
if line.startswith("-") or re.match(r'^\d+\.', line):
|
||||
prompts.append(line[2:])
|
||||
|
||||
gpt_concept_name = extract_gpt_concept_name(gpt_completion, concept_mode)
|
||||
|
||||
trigger_text = "TOK, " + gpt_concept_name if concept_mode == 'object_injection' else "TOK"
|
||||
|
||||
gpt_concept_description = extract_gpt_concept_description(gpt_completion, concept_mode)
|
||||
trigger_text = "TOK"
|
||||
if concept_mode == 'style':
|
||||
trigger_text = ", in the style of TOK"
|
||||
gpt_concept_name = "" # Disables segmentation for style (use full img)
|
||||
trigger_text = "in the style of TOK, "
|
||||
|
||||
return prompts, gpt_concept_name, trigger_text
|
||||
return prompts, gpt_concept_description, trigger_text
|
||||
|
||||
def extract_gpt_concept_name(gpt_completion, concept_mode):
|
||||
def extract_gpt_concept_description(gpt_completion, concept_mode):
|
||||
"""
|
||||
Extracts the concept name from the GPT completion based on the concept mode.
|
||||
"""
|
||||
concept_name = ""
|
||||
prefix = ""
|
||||
if concept_mode in ['face', 'style']:
|
||||
concept_name = concept_mode
|
||||
prefix = "Style Name:" if concept_mode == 'style' else ""
|
||||
elif concept_mode in ['object_injection', 'object']:
|
||||
prefix = "Concept Name:"
|
||||
concept_mode = 'object_injection'
|
||||
|
||||
if prefix:
|
||||
for line in gpt_completion.split("\n"):
|
||||
if line.startswith(prefix):
|
||||
concept_name = line[len(prefix):].strip()
|
||||
break
|
||||
if concept_mode == 'face':
|
||||
prefix = "TOK Description:"
|
||||
elif concept_mode == 'style':
|
||||
prefix = "Style Description:"
|
||||
elif concept_mode == 'object':
|
||||
prefix = "Concept Description:"
|
||||
|
||||
for line in gpt_completion.split("\n"):
|
||||
if line.startswith(prefix):
|
||||
concept_name = line[len(prefix):].strip()
|
||||
break
|
||||
|
||||
return concept_name
|
||||
|
||||
|
||||
def post_process_captions(captions, text, concept_mode, job_seed):
|
||||
|
||||
text = text.strip()
|
||||
print(f"Input captioning text: {text}")
|
||||
gpt_cleanup_worked = False
|
||||
gpt_concept_description = None
|
||||
|
||||
if len(captions) > 3 and len(captions) < MAX_GPT_PROMPTS and not text and client:
|
||||
if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client:
|
||||
retry_count = 0
|
||||
while retry_count < 10:
|
||||
while retry_count < 5:
|
||||
try:
|
||||
gpt_captions, gpt_concept_name, trigger_text = cleanup_prompts_with_chatgpt(captions, concept_mode, job_seed + retry_count)
|
||||
gpt_captions, gpt_concept_description, trigger_text = cleanup_prompts_with_chatgpt(captions, concept_mode, job_seed + retry_count)
|
||||
n_toks = sum("TOK" in caption for caption in gpt_captions)
|
||||
|
||||
if n_toks > int(0.8 * len(captions)) and (len(gpt_captions) == len(captions)):
|
||||
# Ensure every caption contains "TOK"
|
||||
# gpt-cleanup (mostly) worked, lets just ensure every caption contains "TOK" and finish
|
||||
print("Making sure TOK is added to every training prompt...")
|
||||
gpt_captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in gpt_captions]
|
||||
captions = gpt_captions
|
||||
gpt_cleanup_worked = True
|
||||
break
|
||||
else:
|
||||
if len(gpt_captions) == len(captions):
|
||||
print(f'GPT-4 did not return enough {n_toks}/{len(captions)} prompts containing "TOK", retrying...')
|
||||
else:
|
||||
print(f'GPT-4 returned the wrong number of prompts {len(gpt_captions)} instead of {len(captions)}, retrying...')
|
||||
retry_count += 1
|
||||
gpt_cleanup_worked = False
|
||||
|
||||
except Exception as e:
|
||||
retry_count += 1
|
||||
gpt_cleanup_worked = False
|
||||
print(f"An error occurred after try {retry_count}: {e}")
|
||||
time.sleep(1)
|
||||
else:
|
||||
gpt_concept_name, trigger_text = None, "TOK"
|
||||
else:
|
||||
time.sleep(0.5)
|
||||
|
||||
if not gpt_cleanup_worked:
|
||||
# simple concat of trigger text with rest of prompt:
|
||||
if len(text) == 0:
|
||||
print("WARNING: no captioning text was given and we're not doing chatgpt cleanup...")
|
||||
@@ -395,16 +374,14 @@ def post_process_captions(captions, text, concept_mode, job_seed):
|
||||
trigger_text = "in the style of TOK, "
|
||||
captions = [trigger_text + caption for caption in captions]
|
||||
else:
|
||||
trigger_text = "a photo of TOK, "
|
||||
trigger_text = "TOK, "
|
||||
captions = [trigger_text + caption for caption in captions]
|
||||
else:
|
||||
trigger_text = text
|
||||
captions = [trigger_text + ", " + caption for caption in captions]
|
||||
|
||||
gpt_concept_name = None
|
||||
|
||||
captions = [fix_prompt(caption) for caption in captions]
|
||||
return captions, trigger_text, gpt_concept_name
|
||||
return captions, trigger_text, gpt_concept_description
|
||||
|
||||
|
||||
def blip_caption_dataset(
|
||||
@@ -416,19 +393,24 @@ def blip_caption_dataset(
|
||||
"Salesforce/blip2-opt-2.7b",
|
||||
] = "Salesforce/blip-image-captioning-large"
|
||||
):
|
||||
|
||||
# If non of the captions are None, we dont need to do anything:
|
||||
if all(captions):
|
||||
print(f"All captions are already generated, skipping captioning...")
|
||||
return captions
|
||||
|
||||
print(f"Using model {model_id} for image captioning...")
|
||||
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
if "blip2" in model_id:
|
||||
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
||||
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
|
||||
model = Blip2ForConditionalGeneration.from_pretrained(
|
||||
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
|
||||
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
|
||||
).to(device)
|
||||
else:
|
||||
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
||||
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
|
||||
model = BlipForConditionalGeneration.from_pretrained(
|
||||
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
|
||||
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
|
||||
).to(device)
|
||||
|
||||
for i, image in enumerate(tqdm(images)):
|
||||
@@ -437,6 +419,10 @@ def blip_caption_dataset(
|
||||
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65)
|
||||
captions[i] = processor.decode(out[0], skip_special_tokens=True)
|
||||
|
||||
del model
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return captions
|
||||
|
||||
def encode_image(image_path):
|
||||
@@ -454,6 +440,51 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
|
||||
os.remove(output_path)
|
||||
return base64_image
|
||||
|
||||
def gpt4_v_get_description(config, images):
|
||||
if config.concept_mode == "object":
|
||||
description = "object"
|
||||
prompt = "Give a concise visual descriptioni of the object/figure/thing that all the grid-images have in common with at most 10 words. Dont start with statements like 'The image features...', just describe what you see."
|
||||
|
||||
elif config.concept_mode == "face":
|
||||
description = "face"
|
||||
prompt = "All the grid images depict a single person. Visually describe this person with at most 10 words. Dont start with statements like 'The image features...', just describe what you see. (eg an asian woman with long black hair)"
|
||||
|
||||
elif config.concept_mode == "style":
|
||||
description = ""
|
||||
prompt = "All these images share a common aesthetic style. Describe this style with at most 7 words. Dont start with statements like 'The image features...', just describe what you see. (eg impressionism collage surrealism)"
|
||||
|
||||
if not OPENAI_API_KEY:
|
||||
print(f"Skipping GPT-4 Vision description because OPENAI_API_KEY is not set.")
|
||||
return description
|
||||
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {OPENAI_API_KEY}"
|
||||
}
|
||||
|
||||
# TODO sample a grid img:
|
||||
# .... TODO
|
||||
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
|
||||
|
||||
payload = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": prompt},
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{base64_image}", "detail": "high"}}
|
||||
]
|
||||
}
|
||||
],
|
||||
"max_tokens": 60
|
||||
}
|
||||
|
||||
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
|
||||
answer = response.json()["choices"][0]["message"]["content"]
|
||||
return captions
|
||||
|
||||
def gpt4_v_caption_dataset(
|
||||
images, captions,
|
||||
batch_size=4,
|
||||
@@ -463,7 +494,7 @@ def gpt4_v_caption_dataset(
|
||||
print(f"Skipping GPT-4 Vision captioning because OPENAI_API_KEY is not set.")
|
||||
return captions
|
||||
|
||||
prompt = "Accurate describe the contents of this image without assumptions. Avoid starting with statements like 'The image features...', just describe what you see."
|
||||
prompt = "Concisely describe this image without assumptions with at most 20 words. Dont start with statements like 'The image features...', just describe what you see."
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
@@ -474,7 +505,7 @@ def gpt4_v_caption_dataset(
|
||||
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
|
||||
|
||||
payload = {
|
||||
"model": "gpt-4-vision-preview",
|
||||
"model": "gpt-4o",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
@@ -484,11 +515,18 @@ def gpt4_v_caption_dataset(
|
||||
]
|
||||
}
|
||||
],
|
||||
"max_tokens": 100
|
||||
"max_tokens": 60
|
||||
}
|
||||
|
||||
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
|
||||
return index, response.json()["choices"][0]["message"]["content"]
|
||||
|
||||
try:
|
||||
result = response.json()["choices"][0]["message"]["content"]
|
||||
except:
|
||||
print(response.json())
|
||||
result = ""
|
||||
|
||||
return index, result
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=batch_size) as executor:
|
||||
future_to_index = {executor.submit(fetch_caption, i, img): i for i, img in enumerate(images) if captions[i] is None}
|
||||
@@ -519,49 +557,6 @@ def caption_dataset(
|
||||
|
||||
return captions
|
||||
|
||||
|
||||
|
||||
def _crop_to_square(
|
||||
image: Image.Image, com: List[Tuple[int, int]], resize_to: Optional[int] = None
|
||||
):
|
||||
cx, cy = com
|
||||
width, height = image.size
|
||||
if width > height:
|
||||
left_possible = max(cx - height / 2, 0)
|
||||
left = min(left_possible, width - height)
|
||||
right = left + height
|
||||
top = 0
|
||||
bottom = height
|
||||
else:
|
||||
left = 0
|
||||
right = width
|
||||
top_possible = max(cy - width / 2, 0)
|
||||
top = min(top_possible, height - width)
|
||||
bottom = top + width
|
||||
|
||||
image = image.crop((left, top, right, bottom))
|
||||
|
||||
if resize_to:
|
||||
image = image.resize((resize_to, resize_to), Image.Resampling.LANCZOS)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
def _center_of_mass(mask: Image.Image):
|
||||
"""
|
||||
Returns the center of mass of the mask
|
||||
"""
|
||||
x, y = np.meshgrid(np.arange(mask.size[0]), np.arange(mask.size[1]))
|
||||
mask_np = np.array(mask) + 0.01
|
||||
x_ = x * mask_np
|
||||
y_ = y * mask_np
|
||||
|
||||
x = np.sum(x_) / np.sum(mask_np)
|
||||
y = np.sum(y_) / np.sum(mask_np)
|
||||
|
||||
return x, y
|
||||
|
||||
|
||||
def load_image_with_orientation(path, mode = "RGB"):
|
||||
image = Image.open(path)
|
||||
|
||||
@@ -640,9 +635,53 @@ def augment_image(image):
|
||||
image = gaussian_blur(image)
|
||||
return image
|
||||
|
||||
def round_to_nearest_multiple(x, multiple):
|
||||
return int(float(multiple) * round(float(x) / float(multiple)))
|
||||
|
||||
'''
|
||||
For Stable Diffusion 1.5, outputs are optimised around 512x512 pixels. Many common fine-tuned versions of SD1.5 are optimised around 768x768. The best resolutions for common aspect ratios are typically:
|
||||
1:1 (square): 512x512, 768x768
|
||||
3:2 (landscape): 768x512
|
||||
2:3 (portrait): 512x768
|
||||
4:3 (landscape): 768x576
|
||||
3:4 (portrait): 576x768
|
||||
16:9 (widescreen): 912x512
|
||||
9:16 (tall): 512x912
|
||||
|
||||
For SDXL, outputs are optimised around 1024x1024 pixels. The best resolutions for common aspect ratios are typically:
|
||||
stable-diffusion-xl-1024-v0-9 supports generating images at the following dimensions:
|
||||
1024 x 1024
|
||||
1152 x 896
|
||||
896 x 1152
|
||||
1216 x 832
|
||||
832 x 1216
|
||||
1344 x 768
|
||||
768 x 1344
|
||||
1536 x 640
|
||||
640 x 1536
|
||||
|
||||
'''
|
||||
|
||||
def calculate_new_dimensions(target_size, target_aspect_ratio):
|
||||
"""
|
||||
Calculate the new width and height given a target size and aspect ratio.
|
||||
"""
|
||||
# Calculate the total number of pixels
|
||||
n_pixels = target_size ** 2
|
||||
|
||||
# Calculate the new width and height based on the target aspect ratio
|
||||
new_width = (n_pixels * target_aspect_ratio) ** 0.5
|
||||
new_height = (n_pixels / new_width)
|
||||
|
||||
# round up/down to the nearest multiple of 64:
|
||||
new_width = round_to_nearest_multiple(new_width, 64)
|
||||
new_height = round_to_nearest_multiple(new_height, 64)
|
||||
|
||||
return [new_width, new_height]
|
||||
|
||||
|
||||
def load_and_save_masks_and_captions(
|
||||
config,
|
||||
concept_mode: str,
|
||||
files: Union[str, List[str]],
|
||||
output_dir: str = "tmp_out",
|
||||
@@ -652,7 +691,6 @@ def load_and_save_masks_and_captions(
|
||||
target_size: int = 1024,
|
||||
crop_based_on_salience: bool = True,
|
||||
use_face_detection_instead: bool = False,
|
||||
temp: float = 1.0,
|
||||
n_length: int = -1,
|
||||
add_lr_flips: bool = False,
|
||||
augment_imgs_up_to_n: int = 0,
|
||||
@@ -668,7 +706,6 @@ def load_and_save_masks_and_captions(
|
||||
# load images
|
||||
if isinstance(files, str):
|
||||
if os.path.isdir(files):
|
||||
print("Scanning directory for images...")
|
||||
files = (
|
||||
_find_files("*.png", files)
|
||||
+ _find_files("*.jpg", files)
|
||||
@@ -677,7 +714,7 @@ def load_and_save_masks_and_captions(
|
||||
|
||||
if len(files) == 0:
|
||||
raise Exception(
|
||||
f"No files found in {files}. Either {files} is not a directory or it does not contain any .png or .jpg/jpeg files."
|
||||
f"No images were found... Are you sure you provided a valid dataset?"
|
||||
)
|
||||
if n_length == -1:
|
||||
n_length = len(files)
|
||||
@@ -693,6 +730,30 @@ def load_and_save_masks_and_captions(
|
||||
else:
|
||||
captions.append(None)
|
||||
|
||||
# Compute average aspect ratio of images:
|
||||
aspect_ratios = [image.size[0] / image.size[1] for image in images]
|
||||
avg_aspect_ratio = sum(aspect_ratios) / len(aspect_ratios)
|
||||
print(f"Average aspect ratio of images (width / height): {avg_aspect_ratio:.3f}")
|
||||
config.train_img_size = calculate_new_dimensions(target_size, avg_aspect_ratio)
|
||||
config.train_aspect_ratio = config.train_img_size[0] / config.train_img_size[1]
|
||||
target_size = max(config.train_img_size)
|
||||
print(f"New train_img_size: {config.train_img_size}")
|
||||
|
||||
if config.validation_img_size is None:
|
||||
config.validation_img_size = [0, 0]
|
||||
multiplier = 2.0 if config.sd_model_version == "sdxl" else 1.0
|
||||
config.validation_img_size[0] = config.train_img_size[0] * multiplier
|
||||
config.validation_img_size[1] = config.train_img_size[1] * multiplier
|
||||
elif isinstance(config.validation_img_size, int):
|
||||
n_pixels = config.validation_img_size ** 2
|
||||
config.validation_img_size = [0, 0]
|
||||
config.validation_img_size[0] = (n_pixels * config.train_aspect_ratio) ** 0.5
|
||||
config.validation_img_size[1] = (n_pixels / config.validation_img_size[0])
|
||||
|
||||
config.validation_img_size[0] = round_to_nearest_multiple(config.validation_img_size[0], 64)
|
||||
config.validation_img_size[1] = round_to_nearest_multiple(config.validation_img_size[1], 64)
|
||||
print(f"Validation_img_size was set to: {config.validation_img_size}")
|
||||
|
||||
n_training_imgs = len(images)
|
||||
n_captions = len([c for c in captions if c is not None])
|
||||
print(f"Loaded {n_training_imgs} images, {n_captions} of which have captions.")
|
||||
@@ -700,27 +761,47 @@ def load_and_save_masks_and_captions(
|
||||
if len(images) < 50: # upscale images that are smaller than target_size:
|
||||
print("upscaling imgs..")
|
||||
upscale_margin = 0.75
|
||||
images = swin_ir_sr(images, target_size=(int(target_size*upscale_margin), int(target_size*upscale_margin)))
|
||||
images = swin_ir_sr(images, target_size=(int(config.train_img_size[0]*upscale_margin), int(config.train_img_size[0]*upscale_margin)))
|
||||
|
||||
if add_lr_flips and len(images) < 40:
|
||||
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
|
||||
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
|
||||
captions = captions + captions
|
||||
|
||||
# It's nice if we can achieve the gpt pass, so pre-augment the images if there's very few:
|
||||
# Ensure we have at least 'augment_imgs_up_to_n' images through augmentation
|
||||
aug_imgs, aug_caps = [],[]
|
||||
# if we still have a very small amount of imgs, do some basic augmentation:
|
||||
while len(images) + len(aug_imgs) < MIN_GPT_PROMPTS:
|
||||
print(f"Adding augmented version of each training img...")
|
||||
aug_imgs.extend([augment_image(image) for image in images])
|
||||
aug_caps.extend(captions)
|
||||
|
||||
images.extend(aug_imgs)
|
||||
captions.extend(aug_caps)
|
||||
|
||||
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
|
||||
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
|
||||
images = images[:MAX_GPT_PROMPTS-1]
|
||||
captions = captions[:MAX_GPT_PROMPTS-1]
|
||||
|
||||
# Use BLIP for autocaptioning:
|
||||
if len(images) > 50 and caption_model != "blip":
|
||||
print(f"Captioning a lot of ({len(images)}) images --> falling back to using blip!")
|
||||
caption_model = "blip"
|
||||
|
||||
print(f"Generating {len(images)} captions using mode: {concept_mode}...")
|
||||
captions = caption_dataset(images, captions, caption_model = caption_model)
|
||||
|
||||
# Cleanup prompts using chatgpt:
|
||||
captions, trigger_text, gpt_concept_name = post_process_captions(captions, caption_text, concept_mode, seed)
|
||||
captions = [fix_prompt(caption) for caption in captions]
|
||||
trigger_text = ""
|
||||
gpt_concept_description = None
|
||||
if not config.disable_ti:
|
||||
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
|
||||
|
||||
aug_imgs, aug_caps = [],[]
|
||||
while len(images) + len(aug_imgs) < augment_imgs_up_to_n: # if we still have a very small amount of imgs, do some basic augmentation:
|
||||
# if we still have a very small amount of imgs, do some basic augmentation:
|
||||
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
|
||||
print(f"Adding augmented version of each training img...")
|
||||
aug_imgs.extend([augment_image(image) for image in images])
|
||||
aug_caps.extend(captions)
|
||||
@@ -728,25 +809,30 @@ def load_and_save_masks_and_captions(
|
||||
images.extend(aug_imgs)
|
||||
captions.extend(aug_caps)
|
||||
|
||||
if (gpt_concept_name is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
|
||||
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_name}")
|
||||
mask_target_prompts = gpt_concept_name
|
||||
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
|
||||
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
|
||||
mask_target_prompts = gpt_concept_description
|
||||
|
||||
if mask_target_prompts is None:
|
||||
if mask_target_prompts is None or config.concept_mode == "style":
|
||||
print("Disabling CLIP-segmentation")
|
||||
mask_target_prompts = ""
|
||||
temp = 999
|
||||
else:
|
||||
temp = config.clipseg_temperature
|
||||
|
||||
print(f"Generating {len(images)} masks...")
|
||||
|
||||
# Make sure we have a bias for the background pixels to never 100% ignore them
|
||||
background_bias = 0.05
|
||||
if not use_face_detection_instead:
|
||||
seg_masks = clipseg_mask_generator(
|
||||
images=images, target_prompts=mask_target_prompts, temp=temp
|
||||
images=images, target_prompts=mask_target_prompts, temp=temp, bias=background_bias
|
||||
)
|
||||
else:
|
||||
mask_target_prompts = "FACE detection was used"
|
||||
if add_lr_flips:
|
||||
print("WARNING you are applying face detection while also doing left-right flips, this might not be what you intended?")
|
||||
seg_masks = face_mask_google_mediapipe(images=images)
|
||||
seg_masks = face_mask_google_mediapipe(images=images, bias=background_bias*255)
|
||||
|
||||
print("Masks generated! Cropping images to center of mass...")
|
||||
# find the center of mass of the mask
|
||||
@@ -756,23 +842,31 @@ def load_and_save_masks_and_captions(
|
||||
coms = [(image.size[0] / 2, image.size[1] / 2) for image in images]
|
||||
|
||||
# based on the center of mass, crop the image to a square
|
||||
print("Cropping squares...")
|
||||
print("Cropping and resizing images...")
|
||||
images = [
|
||||
_crop_to_square(image, com, resize_to=None)
|
||||
_crop_to_aspect_ratio(image, com, target_aspect_ratio = config.train_aspect_ratio, # width / height
|
||||
resize_to = target_size)
|
||||
for image, com in zip(images, coms)
|
||||
]
|
||||
|
||||
seg_masks = [
|
||||
_crop_to_square(mask, com, resize_to=target_size)
|
||||
_crop_to_aspect_ratio(mask, com, target_aspect_ratio = config.train_aspect_ratio, # width / height
|
||||
resize_to = target_size)
|
||||
for mask, com in zip(seg_masks, coms)
|
||||
]
|
||||
|
||||
print("Resizing images to training size...")
|
||||
images = [
|
||||
image.resize((target_size, target_size), Image.Resampling.LANCZOS)
|
||||
for image in images
|
||||
]
|
||||
print("Expanding masks...")
|
||||
if use_face_detection_instead:
|
||||
dilation_radius = -0.02 * (config.train_img_size[0] + config.train_img_size[0]) / 2
|
||||
blur_radius = 0.02 * (config.train_img_size[0] + config.train_img_size[0]) / 2
|
||||
else:
|
||||
dilation_radius = 0.0
|
||||
blur_radius = 0.005 * (config.train_img_size[0] + config.train_img_size[0]) / 2
|
||||
|
||||
for i in range(len(seg_masks)):
|
||||
seg_masks[i] = grow_mask(seg_masks[i], dilation_radius=dilation_radius, blur_radius=blur_radius)
|
||||
print("Done!")
|
||||
|
||||
data = []
|
||||
# clean TEMP_OUT_DIR first
|
||||
if os.path.exists(output_dir):
|
||||
@@ -780,9 +874,22 @@ def load_and_save_masks_and_captions(
|
||||
os.remove(os.path.join(output_dir, file))
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
if config.disable_ti:
|
||||
print('------------------ WARNING -------------------')
|
||||
print("Removing 'TOK, ' from captions...")
|
||||
print("This will completely disable textual_inversion!!")
|
||||
print('------------------ WARNING -------------------')
|
||||
if gpt_concept_description:
|
||||
replace_str = gpt_concept_description
|
||||
else:
|
||||
replace_str = ""
|
||||
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
|
||||
captions = [caption.replace("TOK", replace_str) for caption in captions]
|
||||
else:
|
||||
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
|
||||
|
||||
# Make sure we've correctly inserted the TOK into every caption:
|
||||
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
|
||||
print("Final captions:")
|
||||
for caption in captions:
|
||||
print(caption)
|
||||
|
||||
@@ -806,12 +913,111 @@ def load_and_save_masks_and_captions(
|
||||
df.to_csv(os.path.join(output_dir, "captions.csv"), index=False)
|
||||
print("---> Training data 100% ready to go!")
|
||||
|
||||
return n_training_imgs, trigger_text, mask_target_prompts, captions
|
||||
# do a final prompt cleaning pass to fix weird commas and spaces:
|
||||
captions = [fix_prompt(caption) for caption in captions]
|
||||
|
||||
# Update the training attributes with some info from the pre-processing:
|
||||
config.training_attributes["n_training_imgs"] = n_training_imgs
|
||||
config.training_attributes["trigger_text"] = trigger_text
|
||||
config.training_attributes["segmentation_prompt"] = mask_target_prompts
|
||||
config.training_attributes["gpt_description"] = gpt_concept_description
|
||||
config.training_attributes["captions"] = captions
|
||||
|
||||
return config
|
||||
|
||||
|
||||
from PIL import Image, ImageFilter, ImageChops
|
||||
|
||||
def grow_mask(mask, dilation_radius=5, blur_radius=3):
|
||||
dilation_radius = int(dilation_radius)
|
||||
blur_radius = int(blur_radius)
|
||||
|
||||
# Load the image
|
||||
mask = mask.convert('L') # Ensure it's in grayscale
|
||||
|
||||
# Get the minimum pixel value in the mask:
|
||||
min_mask_value = int(np.min(np.array(mask)))
|
||||
|
||||
# Dilate the mask
|
||||
if dilation_radius > 0:
|
||||
mask = mask.filter(ImageFilter.MinFilter(dilation_radius * 2 + 1))
|
||||
|
||||
# Apply Gaussian blur to the dilated mask
|
||||
if blur_radius > 0:
|
||||
mask = mask.filter(ImageFilter.GaussianBlur(blur_radius))
|
||||
|
||||
# Clip the mask pixel values to make sure they dont go below the minimum value
|
||||
mask = ImageChops.lighter(mask, Image.new('L', mask.size, min_mask_value))
|
||||
|
||||
return mask
|
||||
|
||||
|
||||
def _center_of_mass(mask: Image.Image):
|
||||
"""
|
||||
Returns the center of mass of the mask
|
||||
"""
|
||||
x, y = np.meshgrid(np.arange(mask.size[0]), np.arange(mask.size[1]))
|
||||
mask_np = np.array(mask) + 0.01
|
||||
x_ = x * mask_np
|
||||
y_ = y * mask_np
|
||||
|
||||
x = np.sum(x_) / np.sum(mask_np)
|
||||
y = np.sum(y_) / np.sum(mask_np)
|
||||
|
||||
return x, y
|
||||
|
||||
def _crop_to_aspect_ratio(
|
||||
image: Image.Image,
|
||||
com: List[Tuple[int, int]],
|
||||
target_aspect_ratio: float = 1.0, # width / height
|
||||
resize_to: Optional[int] = None
|
||||
):
|
||||
"""
|
||||
Crops the image to the specified aspect ratio around the center of mass of the mask.
|
||||
"""
|
||||
cx, cy = com
|
||||
width, height = image.size
|
||||
|
||||
if target_aspect_ratio > 1: # Wider than tall
|
||||
new_width = int(min(width, height * target_aspect_ratio))
|
||||
new_height = int(new_width / target_aspect_ratio)
|
||||
else: # Taller than wide or square
|
||||
new_height = int(min(height, width / target_aspect_ratio))
|
||||
new_width = int(new_height * target_aspect_ratio)
|
||||
|
||||
left = int(max(cx - new_width / 2, 0))
|
||||
right = int(min(left + new_width, width))
|
||||
top = int(max(cy - new_height / 2, 0))
|
||||
bottom = int(min(top + new_height, height))
|
||||
|
||||
# Adjust if the crop goes beyond the image boundaries
|
||||
if right > width:
|
||||
overshoot = right - width
|
||||
right = width
|
||||
left = max(0, left - overshoot) # Adjust left as well symmetrically
|
||||
|
||||
if bottom > height:
|
||||
overshoot = bottom - height
|
||||
bottom = height
|
||||
top = max(0, top - overshoot) # Adjust top as well symmetrically
|
||||
|
||||
image = image.crop((left, top, right, bottom))
|
||||
|
||||
if resize_to:
|
||||
if target_aspect_ratio > 1:
|
||||
resize_height = int(resize_to / target_aspect_ratio)
|
||||
image = image.resize((resize_to, resize_height), Image.Resampling.LANCZOS)
|
||||
else:
|
||||
resize_width = int(resize_to * target_aspect_ratio)
|
||||
image = image.resize((resize_width, resize_to), Image.Resampling.LANCZOS)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
|
||||
|
||||
def face_mask_google_mediapipe(
|
||||
images: List[Image.Image], blur_amount: float = 0.0, bias: float = 50.0
|
||||
images: List[Image.Image], blur_amount: float = 0.0, bias: float = 10.0
|
||||
) -> List[Image.Image]:
|
||||
"""
|
||||
Returns a list of images with masks on the face parts.
|
||||
@@ -852,8 +1058,6 @@ def face_mask_google_mediapipe(
|
||||
min(ih - bbox[1], bbox[3]),
|
||||
)
|
||||
|
||||
print(bbox)
|
||||
|
||||
# Extract face landmarks
|
||||
face_landmarks = face_mesh.process(
|
||||
image_np[bbox[1] : bbox[1] + bbox[3], bbox[0] : bbox[0] + bbox[2]]
|
||||
@@ -929,15 +1133,14 @@ def face_mask_google_mediapipe(
|
||||
|
||||
# Convert mask to 'L' mode (grayscale) before saving
|
||||
mask = mask.convert("L")
|
||||
|
||||
masks.append(mask)
|
||||
else:
|
||||
# If face landmarks are not available, add a black mask of the same size as the image
|
||||
masks.append(Image.new("L", (iw, ih), 255))
|
||||
masks.append(Image.new("L", (iw, ih), 0))
|
||||
|
||||
else:
|
||||
print("No face detected, adding full mask")
|
||||
# If no face is detected, add a white mask of the same size as the image
|
||||
masks.append(Image.new("L", (iw, ih), 255))
|
||||
# If no face is detected, add a black mask of the same size as the image
|
||||
masks.append(Image.new("L", (iw, ih), 0))
|
||||
|
||||
return masks
|
||||
@@ -0,0 +1,289 @@
|
||||
from functools import reduce
|
||||
from diffusers import StableDiffusionXLPipeline
|
||||
from diffusers.models.attention_processor import AttnProcessor2_0, Attention
|
||||
from typing import Optional
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.utils.deprecation_utils import deprecate
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
from einops.layers.torch import Reduce
|
||||
from torchtyping import TensorType
|
||||
from einops import rearrange
|
||||
|
||||
# Find all instances of AttnProcessor2_0 in the UNet
|
||||
def find_attnprocessor2_0(unet):
|
||||
|
||||
module_names = []
|
||||
"""
|
||||
this function assumes that there are fewer than 50 down blocks, attention modules and transformer blocks
|
||||
if you're not sure, feel free to set it to an arbitrarily large number
|
||||
don't worry, it won't slow anything down.
|
||||
"""
|
||||
|
||||
for block_type in ["down_blocks", "up_blocks"]:
|
||||
for down_block_index in range(50):
|
||||
for attentions_index in range(50):
|
||||
for transformer_blocks_index in range(50):
|
||||
example_module_name = f"{block_type}.{down_block_index}.attentions.{attentions_index}.transformer_blocks.{transformer_blocks_index}.attn2.processor"
|
||||
|
||||
try:
|
||||
module = get_module_by_name(module=unet, name = example_module_name)
|
||||
assert isinstance(module, AttnProcessor2_0), f"Expected module to be an instance of AttnProcessor2_0 but found it to be: {type(module)}"
|
||||
# print(f"Found: {example_module_name}")
|
||||
module_names.append(example_module_name)
|
||||
except AttributeError:
|
||||
# print(f"Ignored name: {example_module_name}\nsince it does not exist")
|
||||
pass
|
||||
print(f"Found: {len(module_names)} modules")
|
||||
return module_names
|
||||
|
||||
class DAAMLossAttnProcessor2_0:
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||
"""
|
||||
|
||||
def __init__(self, name: str):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
self.name = name
|
||||
self.cross_attention_scores = None
|
||||
self.reduce_op = Reduce(
|
||||
"batch heads img text -> batch img text",
|
||||
reduction="sum"
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
temb: Optional[torch.Tensor] = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if len(args) > 0 or kwargs.get("scale", None) is not None:
|
||||
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
|
||||
deprecate("scale", "1.0.0", deprecation_message)
|
||||
|
||||
residual = hidden_states
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
"""
|
||||
Mayukh's experiment
|
||||
"""
|
||||
mayukh_experiment = False
|
||||
if encoder_hidden_states is not None:
|
||||
"""
|
||||
this triggers cross attn
|
||||
"""
|
||||
mayukh_experiment = True
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
if mayukh_experiment:
|
||||
# Calculate QK^T
|
||||
qk_t = torch.matmul(query, key.transpose(-2, -1))
|
||||
|
||||
# Calculate attention scores (scaled QK^T)
|
||||
d_k = query.size(-1) # Assuming the last dimension is the embedding dimension
|
||||
attention_scores = qk_t / math.sqrt(d_k)
|
||||
|
||||
attention_scores = self.reduce_op(
|
||||
attention_scores,
|
||||
)
|
||||
self.cross_attention_scores = attention_scores
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
class DAAMLoss:
|
||||
def __init__(self, attention_processors: list[DAAMLossAttnProcessor2_0]):
|
||||
self.attention_processors = attention_processors
|
||||
self.layer_names = [
|
||||
x.name for x in attention_processors
|
||||
]
|
||||
|
||||
def get_all_cross_attention_scores(self):
|
||||
cross_attention_scores = {}
|
||||
|
||||
for p in self.attention_processors:
|
||||
cross_attention_scores[
|
||||
p.name
|
||||
] = p.cross_attention_scores
|
||||
|
||||
return cross_attention_scores
|
||||
|
||||
def compute_single_token_loss(self, text_token_index: list[int], reduce = False):
|
||||
loss = {}
|
||||
cross_attention_scores = self.get_all_cross_attention_scores()
|
||||
|
||||
for name, cross_attention_map in cross_attention_scores.items():
|
||||
"""
|
||||
cross_attention_map.shape: (batch, image_patches, text_tokens)
|
||||
"""
|
||||
assert cross_attention_map.ndim == 3
|
||||
loss[name] = cross_attention_map[:,:,text_token_index].norm() / cross_attention_map.shape[1]
|
||||
|
||||
if reduce:
|
||||
all_losses = list(loss.values())
|
||||
return sum(all_losses)/len(all_losses)
|
||||
else:
|
||||
return loss
|
||||
|
||||
def compute_loss(self, text_token_indices: list[int], reduce = False):
|
||||
losses = []
|
||||
for text_token_index in text_token_indices:
|
||||
losses.append(
|
||||
self.compute_single_token_loss(
|
||||
text_token_index=text_token_index,
|
||||
reduce = True
|
||||
)
|
||||
)
|
||||
|
||||
if reduce:
|
||||
return sum(losses)/len(losses)
|
||||
else:
|
||||
return losses
|
||||
|
||||
def get_image_heatmap(self, text_token_index: int, layer_name: str) -> TensorType["batch", "height", "width"]:
|
||||
cross_attention_scores = self.get_all_cross_attention_scores()
|
||||
assert layer_name in list(cross_attention_scores.keys())
|
||||
|
||||
cross_attention_scores_single_token = cross_attention_scores[layer_name][:,:,text_token_index]
|
||||
assert cross_attention_scores_single_token.ndim == 2 ## batch, hw
|
||||
|
||||
heatmap = rearrange(
|
||||
cross_attention_scores_single_token,
|
||||
"batch (height width) -> batch height width",
|
||||
height = int(math.sqrt(cross_attention_scores_single_token.shape[1])),
|
||||
width = int(math.sqrt(cross_attention_scores_single_token.shape[1]))
|
||||
)
|
||||
|
||||
return heatmap
|
||||
|
||||
def get_the_daam_heatmap(self, text_token_index: int) ->TensorType["batch", "height", "width"]:
|
||||
all_heatmaps = []
|
||||
for layer_name in self.layer_names:
|
||||
heatmap = self.get_image_heatmap(
|
||||
text_token_index=text_token_index,
|
||||
layer_name=layer_name
|
||||
)
|
||||
all_heatmaps.append(heatmap)
|
||||
|
||||
## each heatmap has a shape: batch, h, w where h=w
|
||||
## now find the maximum possible height and width across all heatmaps
|
||||
max_height = max(heatmap.shape[1] for heatmap in all_heatmaps)
|
||||
max_width = max(heatmap.shape[2] for heatmap in all_heatmaps)
|
||||
|
||||
## now resize all_heatmaps to (batch, max_height, max_width) using F.interpolate
|
||||
resized_heatmaps = [
|
||||
F.interpolate(input = x.unsqueeze(1), size = (max_height, max_width)).squeeze(1)
|
||||
for x in all_heatmaps
|
||||
]
|
||||
|
||||
return sum(resized_heatmaps)
|
||||
|
||||
|
||||
def get_module_by_name(module: nn.Module, name: str):
|
||||
"""Retrieve a module nested in another by its access string."""
|
||||
if name == "":
|
||||
return module
|
||||
names = name.split(sep=".")
|
||||
return reduce(getattr, names, module)
|
||||
|
||||
def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionXLPipeline, DAAMLoss]:
|
||||
|
||||
assert isinstance(pipeline, StableDiffusionXLPipeline)
|
||||
|
||||
## find out where the attention processor thingies are
|
||||
module_names = find_attnprocessor2_0(
|
||||
unet = pipeline.unet
|
||||
)
|
||||
|
||||
all_daam_attention_processors = []
|
||||
# override the attention processor thingies
|
||||
for name in module_names:
|
||||
# print(f"Replacing: {name}")
|
||||
|
||||
# Get parent module and attribute name
|
||||
parent_name = ".".join(name.split(".")[:-1])
|
||||
attr_name = name.split(".")[-1]
|
||||
|
||||
# Get the parent module
|
||||
parent_module = get_module_by_name(module=pipeline.unet, name=parent_name)
|
||||
|
||||
daam_attention_processor = DAAMLossAttnProcessor2_0(name=name)
|
||||
all_daam_attention_processors.append(daam_attention_processor)
|
||||
# Set the attribute
|
||||
setattr(parent_module, attr_name, daam_attention_processor)
|
||||
|
||||
# Verify the replacement
|
||||
current_module = get_module_by_name(module=pipeline.unet, name=name)
|
||||
assert isinstance(current_module, DAAMLossAttnProcessor2_0)
|
||||
|
||||
daam_loss = DAAMLoss(
|
||||
attention_processors=all_daam_attention_processors
|
||||
)
|
||||
return pipeline, daam_loss
|
||||
@@ -0,0 +1,267 @@
|
||||
# Released under MIT license
|
||||
# Copyright (c) 2022 finetuneanon (NovelAI/Anlatan LLC)
|
||||
|
||||
import numpy as np
|
||||
import pickle
|
||||
import time
|
||||
|
||||
def get_prng(seed):
|
||||
return np.random.RandomState(seed)
|
||||
|
||||
class BucketManager:
|
||||
def __init__(self, aspect_ratios, valid_ids=None, max_size=(768,512), divisible=64, step_size=8, min_dim=256, base_res=(512,512), bsz=1, world_size=1, global_rank=0, max_ar_error=4, seed=42, dim_limit=2048, debug=False):
|
||||
|
||||
self.res_map = aspect_ratios
|
||||
if valid_ids is not None:
|
||||
new_res_map = {}
|
||||
valid_ids = set(valid_ids)
|
||||
for k, v in self.res_map.items():
|
||||
if k in valid_ids:
|
||||
new_res_map[k] = v
|
||||
self.res_map = new_res_map
|
||||
self.max_size = max_size
|
||||
self.f = 8
|
||||
self.max_tokens = (max_size[0]/self.f) * (max_size[1]/self.f)
|
||||
self.div = divisible
|
||||
self.min_dim = min_dim
|
||||
self.dim_limit = dim_limit
|
||||
self.base_res = base_res
|
||||
self.bsz = bsz
|
||||
self.world_size = world_size
|
||||
self.global_rank = global_rank
|
||||
self.max_ar_error = max_ar_error
|
||||
self.prng = get_prng(seed)
|
||||
epoch_seed = self.prng.tomaxint() % (2**32-1)
|
||||
self.epoch_prng = get_prng(epoch_seed) # separate prng for sharding use for increased thread resilience
|
||||
self.epoch = None
|
||||
self.left_over = None
|
||||
self.batch_total = None
|
||||
self.batch_delivered = None
|
||||
|
||||
self.debug = debug
|
||||
|
||||
self.gen_buckets()
|
||||
self.assign_buckets()
|
||||
self.start_epoch()
|
||||
|
||||
def gen_buckets(self):
|
||||
if self.debug:
|
||||
timer = time.perf_counter()
|
||||
resolutions = []
|
||||
aspects = []
|
||||
w = self.min_dim
|
||||
while (w/self.f) * (self.min_dim/self.f) <= self.max_tokens and w <= self.dim_limit:
|
||||
h = self.min_dim
|
||||
got_base = False
|
||||
while (w/self.f) * ((h+self.div)/self.f) <= self.max_tokens and (h+self.div) <= self.dim_limit:
|
||||
if w == self.base_res[0] and h == self.base_res[1]:
|
||||
got_base = True
|
||||
h += self.div
|
||||
if (w != self.base_res[0] or h != self.base_res[1]) and got_base:
|
||||
resolutions.append(self.base_res)
|
||||
aspects.append(1)
|
||||
resolutions.append((w, h))
|
||||
aspects.append(float(w)/float(h))
|
||||
w += self.div
|
||||
h = self.min_dim
|
||||
while (h/self.f) * (self.min_dim/self.f) <= self.max_tokens and h <= self.dim_limit:
|
||||
w = self.min_dim
|
||||
got_base = False
|
||||
while (h/self.f) * ((w+self.div)/self.f) <= self.max_tokens and (w+self.div) <= self.dim_limit:
|
||||
if w == self.base_res[0] and h == self.base_res[1]:
|
||||
got_base = True
|
||||
w += self.div
|
||||
resolutions.append((w, h))
|
||||
aspects.append(float(w)/float(h))
|
||||
h += self.div
|
||||
res_map = {}
|
||||
for i, res in enumerate(resolutions):
|
||||
res_map[res] = aspects[i]
|
||||
self.resolutions = sorted(res_map.keys(), key=lambda x: x[0] * 4096 - x[1])
|
||||
self.aspects = np.array(list(map(lambda x: res_map[x], self.resolutions)))
|
||||
self.resolutions = np.array(self.resolutions)
|
||||
if self.debug:
|
||||
timer = time.perf_counter() - timer
|
||||
print(f"resolutions:\n{self.resolutions}")
|
||||
print(f"aspects:\n{self.aspects}")
|
||||
print(f"gen_buckets: {timer:.5f}s")
|
||||
|
||||
def assign_buckets(self):
|
||||
if self.debug:
|
||||
timer = time.perf_counter()
|
||||
self.buckets = {}
|
||||
self.aspect_errors = []
|
||||
skipped = 0
|
||||
skip_list = []
|
||||
for post_id in self.res_map.keys():
|
||||
w, h = self.res_map[post_id]
|
||||
aspect = float(w)/float(h)
|
||||
bucket_id = np.abs(self.aspects - aspect).argmin()
|
||||
if bucket_id not in self.buckets:
|
||||
self.buckets[bucket_id] = []
|
||||
error = abs(self.aspects[bucket_id] - aspect)
|
||||
if error < self.max_ar_error:
|
||||
self.buckets[bucket_id].append(post_id)
|
||||
if self.debug:
|
||||
self.aspect_errors.append(error)
|
||||
else:
|
||||
skipped += 1
|
||||
skip_list.append(post_id)
|
||||
for post_id in skip_list:
|
||||
del self.res_map[post_id]
|
||||
if self.debug:
|
||||
timer = time.perf_counter() - timer
|
||||
self.aspect_errors = np.array(self.aspect_errors)
|
||||
print(f"skipped images: {skipped}")
|
||||
print(f"aspect error: mean {self.aspect_errors.mean()}, median {np.median(self.aspect_errors)}, max {self.aspect_errors.max()}")
|
||||
for bucket_id in reversed(sorted(self.buckets.keys(), key=lambda b: len(self.buckets[b]))):
|
||||
print(f"bucket {bucket_id}: {self.resolutions[bucket_id]}, aspect {self.aspects[bucket_id]:.5f}, entries {len(self.buckets[bucket_id])}")
|
||||
print(f"assign_buckets: {timer:.5f}s")
|
||||
|
||||
def start_epoch(self, world_size=None, global_rank=None):
|
||||
if self.debug:
|
||||
timer = time.perf_counter()
|
||||
if world_size is not None:
|
||||
self.world_size = world_size
|
||||
if global_rank is not None:
|
||||
self.global_rank = global_rank
|
||||
|
||||
# select ids for this epoch/rank
|
||||
index = np.array(sorted(list(self.res_map.keys())))
|
||||
index_len = index.shape[0]
|
||||
index = self.epoch_prng.permutation(index)
|
||||
index = index[:index_len - (index_len % (self.bsz * self.world_size))]
|
||||
#print("perm", self.global_rank, index[0:16])
|
||||
index = index[self.global_rank::self.world_size]
|
||||
self.batch_total = index.shape[0] // self.bsz
|
||||
assert(index.shape[0] % self.bsz == 0)
|
||||
index = set(index)
|
||||
|
||||
self.epoch = {}
|
||||
self.left_over = []
|
||||
self.batch_delivered = 0
|
||||
for bucket_id in sorted(self.buckets.keys()):
|
||||
if len(self.buckets[bucket_id]) > 0:
|
||||
self.epoch[bucket_id] = np.array([post_id for post_id in self.buckets[bucket_id] if post_id in index], dtype=np.int64)
|
||||
self.prng.shuffle(self.epoch[bucket_id])
|
||||
self.epoch[bucket_id] = list(self.epoch[bucket_id])
|
||||
overhang = len(self.epoch[bucket_id]) % self.bsz
|
||||
if overhang != 0:
|
||||
self.left_over.extend(self.epoch[bucket_id][:overhang])
|
||||
self.epoch[bucket_id] = self.epoch[bucket_id][overhang:]
|
||||
if len(self.epoch[bucket_id]) == 0:
|
||||
del self.epoch[bucket_id]
|
||||
|
||||
if self.debug:
|
||||
timer = time.perf_counter() - timer
|
||||
count = 0
|
||||
for bucket_id in self.epoch.keys():
|
||||
count += len(self.epoch[bucket_id])
|
||||
print(f"correct item count: {count == len(index)} ({count} of {len(index)})")
|
||||
print(f"start_epoch: {timer:.5f}s")
|
||||
|
||||
def get_batch(self):
|
||||
if self.debug:
|
||||
timer = time.perf_counter()
|
||||
# check if no data left or no epoch initialized
|
||||
if self.epoch is None or self.left_over is None or (len(self.left_over) == 0 and not bool(self.epoch)) or self.batch_total == self.batch_delivered:
|
||||
self.start_epoch()
|
||||
|
||||
found_batch = False
|
||||
batch_data = None
|
||||
resolution = self.base_res
|
||||
while not found_batch:
|
||||
bucket_ids = list(self.epoch.keys())
|
||||
if len(self.left_over) >= self.bsz:
|
||||
bucket_probs = [len(self.left_over)] + [len(self.epoch[bucket_id]) for bucket_id in bucket_ids]
|
||||
bucket_ids = [-1] + bucket_ids
|
||||
else:
|
||||
bucket_probs = [len(self.epoch[bucket_id]) for bucket_id in bucket_ids]
|
||||
bucket_probs = np.array(bucket_probs, dtype=np.float32)
|
||||
bucket_lens = bucket_probs
|
||||
bucket_probs = bucket_probs / bucket_probs.sum()
|
||||
bucket_ids = np.array(bucket_ids, dtype=np.int64)
|
||||
if bool(self.epoch):
|
||||
chosen_id = int(self.prng.choice(bucket_ids, 1, p=bucket_probs)[0])
|
||||
else:
|
||||
chosen_id = -1
|
||||
|
||||
if chosen_id == -1:
|
||||
# using leftover images that couldn't make it into a bucketed batch and returning them for use with basic square image
|
||||
self.prng.shuffle(self.left_over)
|
||||
batch_data = self.left_over[:self.bsz]
|
||||
self.left_over = self.left_over[self.bsz:]
|
||||
found_batch = True
|
||||
else:
|
||||
if len(self.epoch[chosen_id]) >= self.bsz:
|
||||
# return bucket batch and resolution
|
||||
batch_data = self.epoch[chosen_id][:self.bsz]
|
||||
self.epoch[chosen_id] = self.epoch[chosen_id][self.bsz:]
|
||||
resolution = tuple(self.resolutions[chosen_id])
|
||||
found_batch = True
|
||||
if len(self.epoch[chosen_id]) == 0:
|
||||
del self.epoch[chosen_id]
|
||||
else:
|
||||
# can't make a batch from this, not enough images. move them to leftovers and try again
|
||||
self.left_over.extend(self.epoch[chosen_id])
|
||||
del self.epoch[chosen_id]
|
||||
|
||||
assert(found_batch or len(self.left_over) >= self.bsz or bool(self.epoch))
|
||||
|
||||
if self.debug:
|
||||
timer = time.perf_counter() - timer
|
||||
print(f"bucket probs: " + ", ".join(map(lambda x: f"{x:.2f}", list(bucket_probs*100))))
|
||||
print(f"chosen id: {chosen_id}")
|
||||
print(f"batch data: {batch_data}")
|
||||
print(f"resolution: {resolution}")
|
||||
print(f"get_batch: {timer:.5f}s")
|
||||
|
||||
self.batch_delivered += 1
|
||||
return (batch_data, resolution)
|
||||
|
||||
def generator(self):
|
||||
if self.batch_delivered >= self.batch_total:
|
||||
self.start_epoch()
|
||||
while self.batch_delivered < self.batch_total:
|
||||
yield self.get_batch()
|
||||
|
||||
if __name__ == "__main__":
|
||||
# prepare a pickle with mapping of dataset IDs to resolutions called resolutions.pkl to use this
|
||||
with open("resolutions.pkl", "rb") as fh:
|
||||
ids = list(pickle.load(fh).keys())
|
||||
|
||||
counts = np.zeros((len(ids),)).astype(np.int64)
|
||||
id_map = {}
|
||||
for i, post_id in enumerate(ids):
|
||||
id_map[post_id] = i
|
||||
|
||||
bm = BucketManager("resolutions.pkl", debug=True, bsz=8, world_size=8, global_rank=3)
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
print("got: " + str(bm.get_batch()))
|
||||
|
||||
bm = BucketManager("resolutions.pkl", bsz=8, world_size=1, global_rank=0, valid_ids=ids[0:16])
|
||||
for _ in range(16):
|
||||
bm.get_batch()
|
||||
print("got from future epoch: " + str(bm.get_batch()))
|
||||
|
||||
bms = []
|
||||
for rank in range(16):
|
||||
bm = BucketManager("resolutions.pkl", bsz=8, world_size=16, global_rank=rank)
|
||||
bms.append(bm)
|
||||
for epoch in range(5):
|
||||
print(f"epoch {epoch}")
|
||||
for i, bm in enumerate(bms):
|
||||
print(f"bm {i}")
|
||||
first = True
|
||||
for ids, res in bm.generator():
|
||||
if first and i == 0:
|
||||
#print(ids)
|
||||
first = False
|
||||
for post_id in ids:
|
||||
counts[id_map[post_id]] += 1
|
||||
print(np.bincount(counts))
|
||||
@@ -9,42 +9,6 @@ import signal
|
||||
import time
|
||||
import numpy as np
|
||||
|
||||
SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
|
||||
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
|
||||
|
||||
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
|
||||
# TODO point this url to the correct full folder structure containing the CLIP text-encoder (this wont actually work rn)
|
||||
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
|
||||
|
||||
# Define model paths and URLs in a dictionary
|
||||
MODEL_INFO = {
|
||||
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL},
|
||||
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL}
|
||||
}
|
||||
|
||||
def download_weights(url, dest):
|
||||
start = time.time()
|
||||
print("downloading url: ", url)
|
||||
print("downloading to: ", dest, '...')
|
||||
|
||||
# Make sure the destination directory exists
|
||||
dest_dir = os.path.dirname(dest)
|
||||
if not os.path.exists(dest_dir):
|
||||
os.makedirs(dest_dir)
|
||||
|
||||
try:
|
||||
subprocess.check_call(["wget", "-q", "-O", dest, url])
|
||||
except subprocess.CalledProcessError as e:
|
||||
print("Error occurred while downloading:")
|
||||
print("Exit status:", e.returncode)
|
||||
print("Output:", e.output)
|
||||
except Exception as e:
|
||||
print("An unexpected error occurred:", e)
|
||||
|
||||
print(f"Downloading {url} took {time.time() - start} seconds")
|
||||
|
||||
|
||||
|
||||
def clean_filename(filename):
|
||||
allowed_chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-_"
|
||||
return ''.join(c for c in filename if c in allowed_chars)
|
||||
@@ -132,11 +96,11 @@ def merge_datasets(path_A, path_B, out_path, token_names):
|
||||
|
||||
|
||||
|
||||
def make_validation_img_grid(img_folder):
|
||||
def make_validation_img_grid(img_folder, rows = 2):
|
||||
"""
|
||||
|
||||
find all the .jpg imgs in img_folder (template = *.jpg)
|
||||
if >=4 validation imgs, create a 2x2 grid of them
|
||||
if >=4 validation imgs, create a rows x n grid of them
|
||||
otherwise just return the first validation img
|
||||
|
||||
"""
|
||||
@@ -148,18 +112,21 @@ def make_validation_img_grid(img_folder):
|
||||
# If less than 4 validation images, return path of the first one
|
||||
return os.path.join(img_folder, validation_imgs[0])
|
||||
else:
|
||||
# If >= 4 validation images, create 2x2 grid
|
||||
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:4]]
|
||||
# If >= 4 validation images, create 2xn grid
|
||||
n_imgs = len(validation_imgs) // rows * rows
|
||||
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:n_imgs]]
|
||||
|
||||
n_cols = int(n_imgs / rows)
|
||||
|
||||
# Assuming all images are the same size, get dimensions of first image
|
||||
width, height = imgs[0].size
|
||||
|
||||
# Create an empty image with 2x2 grid size
|
||||
grid_img = Image.new("RGB", (2 * width, 2 * height))
|
||||
grid_img = Image.new("RGB", (n_cols * width, rows * height))
|
||||
|
||||
# Paste the images into the grid
|
||||
for i in range(2):
|
||||
for j in range(2):
|
||||
for i in range(n_cols):
|
||||
for j in range(rows):
|
||||
grid_img.paste(imgs.pop(0), (i * width, j * height))
|
||||
|
||||
# Save the new image
|
||||
@@ -415,7 +382,7 @@ def download_and_prep_training_data(data_location, data_dir):
|
||||
print("Downloading training data...")
|
||||
# we're assuming the data is prived as pipe seperated urls to .zip files
|
||||
for url in str(data_location).split('|'):
|
||||
download(url, data_dir)
|
||||
download(url.strip(), data_dir)
|
||||
|
||||
# Loop over all files in the data directory:
|
||||
for filename in os.listdir(data_dir):
|
||||
@@ -0,0 +1,14 @@
|
||||
import ujson
|
||||
import os
|
||||
|
||||
|
||||
def save_as_json(dictionary_or_list, filename: str):
|
||||
with open(filename, "w") as fp:
|
||||
ujson.dump(dictionary_or_list, fp, indent=4)
|
||||
|
||||
|
||||
def load_json(filename: str):
|
||||
assert os.path.exists(filename), f"Could not find json file: {filename}"
|
||||
with open(filename) as json_file:
|
||||
data = ujson.load(json_file)
|
||||
return data
|
||||
Executable
+280
@@ -0,0 +1,280 @@
|
||||
import os
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import random
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import gc
|
||||
import PIL
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from PIL import Image
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import save_file
|
||||
from torch.utils.data import Dataset
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
import torch.nn.functional as F
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
dtype_map = {
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
"fp32": torch.float32
|
||||
}
|
||||
|
||||
import re
|
||||
def replace_in_string(s, replacements):
|
||||
while True:
|
||||
replaced = False
|
||||
for target, replacement in replacements.items():
|
||||
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
|
||||
if new_s != s:
|
||||
s = new_s
|
||||
replaced = True
|
||||
if not replaced:
|
||||
break
|
||||
return s
|
||||
|
||||
def fix_prompt(prompt: str):
|
||||
if not prompt:
|
||||
return prompt
|
||||
# Remove extra commas and spaces, and fix space before punctuation
|
||||
prompt = re.sub(r"\s+", " ", prompt) # Replace multiple spaces with a single space
|
||||
prompt = re.sub(r",,", ",", prompt) # Replace double commas with a single comma
|
||||
prompt = re.sub(r"\s?,\s?", ", ", prompt) # Fix spaces around commas
|
||||
prompt = re.sub(r"\s?\.\s?", ". ", prompt) # Fix spaces around periods
|
||||
return prompt.strip() # Remove leading and trailing whitespace
|
||||
|
||||
def seed_everything(seed: int):
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
def zipdir(path, ziph, extension = '.py'):
|
||||
# Zip the directory
|
||||
for root, dirs, files in os.walk(path):
|
||||
for file in files:
|
||||
if file.endswith(extension):
|
||||
ziph.write(os.path.join(root, file),
|
||||
os.path.relpath(os.path.join(root, file),
|
||||
os.path.join(path, '..')))
|
||||
|
||||
def pick_best_gpu_id():
|
||||
try:
|
||||
# pick the GPU with the most free memory:
|
||||
gpu_ids = [i for i in range(torch.cuda.device_count())]
|
||||
print(f"# of visible GPUs: {len(gpu_ids)}")
|
||||
gpu_mem = []
|
||||
for gpu_id in gpu_ids:
|
||||
free_memory, tot_mem = torch.cuda.mem_get_info(device=gpu_id)
|
||||
gpu_mem.append(free_memory)
|
||||
print("GPU %d: %d MB free" %(gpu_id, free_memory / 1024 / 1024))
|
||||
|
||||
if len(gpu_ids) == 0:
|
||||
# no GPUs available, use CPU:
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
return None
|
||||
|
||||
best_gpu_id = gpu_ids[np.argmax(gpu_mem)]
|
||||
# set this to be the active GPU:
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(best_gpu_id)
|
||||
print("Using GPU %d" %best_gpu_id)
|
||||
return best_gpu_id
|
||||
except Exception as e:
|
||||
print(f'Error picking best gpu: {e}')
|
||||
print(f'Falling back to GPU 0')
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
|
||||
return 0
|
||||
|
||||
|
||||
import psutil
|
||||
def print_system_info():
|
||||
try:
|
||||
# Print GPU memory information
|
||||
gpu_ids = [i for i in range(torch.cuda.device_count())]
|
||||
for gpu_id in gpu_ids:
|
||||
free_memory, total_memory = torch.cuda.mem_get_info(device=gpu_id)
|
||||
print(f"GPU {gpu_id}: {free_memory // 1024 // 1024} of {total_memory // 1024 // 1024} Mb free")
|
||||
|
||||
# Print disk space information
|
||||
disk_usage = psutil.disk_usage('/')
|
||||
total_disk = disk_usage.total // (1024 * 1024)
|
||||
used_disk = disk_usage.used // (1024 * 1024)
|
||||
percent_disk_used = disk_usage.percent
|
||||
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
|
||||
|
||||
# Print RAM information
|
||||
virtual_mem = psutil.virtual_memory()
|
||||
total_ram = virtual_mem.total // (1024 * 1024)
|
||||
current_ram = virtual_mem.used // (1024 * 1024)
|
||||
percent_ram_used = virtual_mem.percent
|
||||
print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
|
||||
|
||||
except Exception as e:
|
||||
print(f'Error in gathering system info: {str(e)}')
|
||||
|
||||
return
|
||||
|
||||
|
||||
def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1, max_val=1, ymax_f = 0.75, color = 'blue'):
|
||||
try:
|
||||
os.makedirs(checkpoint_dir, exist_ok=True)
|
||||
|
||||
# Flatten and concatenate all parameters into a single tensor
|
||||
all_params = torch.cat([p.data.view(-1) for p in parameters])
|
||||
|
||||
# count number of parameters:
|
||||
n_params = len(all_params)
|
||||
|
||||
if n_params == 0 or n_params > 1e9:
|
||||
return
|
||||
|
||||
norm = torch.norm(all_params)
|
||||
|
||||
# Convert to CPU for plotting
|
||||
all_params_cpu = all_params.cpu().float().numpy()
|
||||
|
||||
# Plot histogram
|
||||
plt.figure()
|
||||
plt.hist(all_params_cpu, bins=bins, density=False, color = color)
|
||||
plt.ylim(0, ymax_f * len(all_params_cpu.flatten()))
|
||||
plt.xlim(min_val, max_val)
|
||||
plt.xlabel('Weight Value')
|
||||
plt.ylabel('Count')
|
||||
plt.title(f'{name} (std: {np.std(all_params_cpu):.5f}, norm: {norm:.3f}, step {step:03d})')
|
||||
plt.savefig(f"{checkpoint_dir}/{name}_hist_{step:04d}.png")
|
||||
plt.close()
|
||||
except:
|
||||
print(f'Error plotting {name} histogram')
|
||||
|
||||
def plot_curve(value_dict, xlabel, ylabel, title, save_path, log_scale = False, y_lims = None):
|
||||
plt.figure()
|
||||
for key in value_dict.keys():
|
||||
values = value_dict[key]
|
||||
plt.plot(range(len(values)), values, label=key)
|
||||
|
||||
if log_scale:
|
||||
plt.yscale('log') # Set y-axis to log scale
|
||||
plt.xlabel(xlabel)
|
||||
plt.ylabel(ylabel)
|
||||
if y_lims is not None:
|
||||
plt.ylim(y_lims[0], y_lims[1])
|
||||
plt.title(title)
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
# plot the learning rates:
|
||||
def plot_lrs(learning_rate_dict, save_path='learning_rates.png'):
|
||||
plt.figure()
|
||||
for key in learning_rate_dict.keys():
|
||||
lrs = learning_rate_dict[key]
|
||||
if len(lrs) == 0:
|
||||
continue
|
||||
plt.plot(range(len(lrs)), lrs, label=key)
|
||||
plt.yscale('log') # Set y-axis to log scale
|
||||
plt.ylim(1e-6, 3e-3)
|
||||
plt.xlabel('Step')
|
||||
plt.ylabel('Learning Rate')
|
||||
plt.title('Learning Rate Curves')
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
# plot the learning rates:
|
||||
def plot_grad_norms(grad_norms, save_path='grad_norms.png'):
|
||||
plt.figure()
|
||||
plt.plot(range(len(grad_norms['unet'])), grad_norms['unet'], label='unet')
|
||||
|
||||
for i in range(2):
|
||||
try:
|
||||
plt.plot(range(len(grad_norms[f'text_encoder_{i}'])), grad_norms[f'text_encoder_{i}'], label=f'text_encoder_{i}')
|
||||
except:
|
||||
pass
|
||||
|
||||
plt.yscale('log') # Set y-axis to log scale
|
||||
plt.ylim(1e-6, 100.0)
|
||||
plt.xlabel('Step')
|
||||
plt.ylabel('Grad Norm')
|
||||
plt.title('Gradient Norms')
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dict = {}):
|
||||
plt.figure()
|
||||
anchor_values = []
|
||||
for key in token_std_dict.keys():
|
||||
tokenizer_i_token_stds = token_std_dict[key]
|
||||
for i in range(len(tokenizer_i_token_stds)):
|
||||
stds = tokenizer_i_token_stds[i]
|
||||
if len(stds) == 0:
|
||||
continue
|
||||
anchor_values.append(stds[0])
|
||||
encoder_index = int(key.split('_')[-1])
|
||||
plt.plot(range(len(stds)), stds, label=f'{key}_tok_{i}', linestyle='dashed' if encoder_index > 0 else 'solid')
|
||||
|
||||
plt.xlabel('Step')
|
||||
plt.ylabel('Token Embedding Std')
|
||||
centre_value = 0.013
|
||||
up_f, down_f = 1.4, 1.3
|
||||
try:
|
||||
plt.ylim(centre_value/down_f, centre_value*up_f)
|
||||
except:
|
||||
pass
|
||||
|
||||
# Plotting target values as horizontal lines
|
||||
for label, value in target_value_dict.items():
|
||||
plt.axhline(y=value, color='r', linestyle='-' if '0' in label else '--', label=label)
|
||||
plt.text(0, value, label, ha='left', va='center')
|
||||
|
||||
plt.title('Token Embedding Std')
|
||||
plt.legend()
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
|
||||
from scipy.signal import savgol_filter
|
||||
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'):
|
||||
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red'}
|
||||
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
|
||||
plot_smoothed = ['img_loss']
|
||||
|
||||
plt.figure(figsize=(8, 5))
|
||||
|
||||
for key, losses in loss_dict.items():
|
||||
if key == 'tot_loss':
|
||||
continue
|
||||
|
||||
losses = np.array(losses)
|
||||
|
||||
if len(losses) < window_length:
|
||||
continue
|
||||
|
||||
if key in plot_smoothed:
|
||||
losses = savgol_filter(losses, window_length, polyorder)
|
||||
label = f'Smoothed {key}'
|
||||
linestyle = 'dashed'
|
||||
else:
|
||||
label = key
|
||||
linestyle = 'solid'
|
||||
|
||||
plot_losses = losses / np.max(losses)
|
||||
color = colormap.get(key, default_color) # Use the default color if the key is not in the colormap
|
||||
plt.plot(plot_losses, label=label, color=color, linestyle=linestyle)
|
||||
|
||||
# Create the title:
|
||||
title = 'Loss values:'
|
||||
for key in values_to_add_to_title:
|
||||
if key in loss_dict:
|
||||
if loss_dict[key]:
|
||||
title += f' {key}: {loss_dict[key][-1]:.3f}'
|
||||
|
||||
plt.title(title)
|
||||
plt.xlabel('Optimizer Step')
|
||||
plt.ylabel('Training Losses')
|
||||
plt.ylim(0, 1.1) # Adjust the y-axis limits for normalized data
|
||||
plt.legend(loc='lower left')
|
||||
plt.savefig(save_path)
|
||||
plt.close()
|
||||
@@ -31,26 +31,21 @@ val_prompts['style'] = [
|
||||
]
|
||||
|
||||
val_prompts["face"] = [
|
||||
"an intricate wood carving of <concept> in a historic temple",
|
||||
'<concept> as pixel art, 8-bit video game style',
|
||||
'painting of <concept> by Vincent van Gogh',
|
||||
'<concept> as a superhero, wearing a cape',
|
||||
'<concept> as a statue made of marble',
|
||||
'<concept> as a character in a noir graphic novel, under a rain-soaked streetlamp',
|
||||
'stop motion animation of <concept> using clay, Wallace and Gromit style',
|
||||
'<concept> portrayed in a famous renaissance painting, replacing Mona Lisas face',
|
||||
'a photo of <concept> attending the Oscars, walking down the red carpet with sunglasses',
|
||||
'<concept> as a pop vinyl figure, complete with oversized head and small body',
|
||||
#'<concept> as a pop vinyl figure, complete with oversized head and small body',
|
||||
'<concept> as a retro holographic sticker, shimmering in bright colors',
|
||||
'<concept> as a bobblehead on a car dashboard, nodding incessantly',
|
||||
"<concept> captured in a snow globe, complete with intricate details",
|
||||
"a photo of <concept> climbing mount Everest in the snow, alpinism",
|
||||
"<concept> as an action figure superhero, lego toy, toy story",
|
||||
'a photo of a massive statue of <concept> in the middle of the city',
|
||||
'a masterful oil painting portraying <concept> with vibrant colors, brushstrokes and textures',
|
||||
'a vibrant low-poly artwork of <concept>, rendered in SVG, vector graphics',
|
||||
'<concept>, polaroid photograph',
|
||||
'a huge <concept> sand sculpture on a sunny beach, made of sand',
|
||||
'an old, vintage, polaroid photograph of <concept>, artsy look',
|
||||
'<concept> immortalized as an exquisite marble statue with masterful chiseling, swirling marble patterns and textures',
|
||||
]
|
||||
|
||||
-783
@@ -1,783 +0,0 @@
|
||||
import fnmatch
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import random
|
||||
import time
|
||||
import shutil
|
||||
import gc
|
||||
import numpy as np
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
import torch.nn.functional as F
|
||||
|
||||
from peft import LoraConfig, get_peft_model
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers import EulerDiscreteScheduler
|
||||
from tqdm import tqdm
|
||||
|
||||
from dataset_and_utils import *
|
||||
from lora_utils import *
|
||||
from io_utils import make_validation_img_grid
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
|
||||
def print_trainable_parameters(model, name = ''):
|
||||
trainable_params = 0
|
||||
all_param = 0
|
||||
for _, param in model.named_parameters():
|
||||
all_param += param.numel()
|
||||
if param.requires_grad:
|
||||
trainable_params += param.numel()
|
||||
line_delimiter = "#" * 70
|
||||
print('\n', line_delimiter)
|
||||
print(
|
||||
f"Trainable {name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
|
||||
)
|
||||
print(line_delimiter, '\n')
|
||||
|
||||
|
||||
def compute_snr(noise_scheduler, timesteps):
|
||||
"""
|
||||
Computes SNR as per
|
||||
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
|
||||
"""
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = alphas_cumprod**0.5
|
||||
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
|
||||
|
||||
# Expand the tensors.
|
||||
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
|
||||
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
|
||||
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
|
||||
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
|
||||
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
|
||||
|
||||
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
|
||||
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
|
||||
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
|
||||
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
|
||||
|
||||
# Compute SNR.
|
||||
snr = (alpha / sigma) ** 2
|
||||
return snr
|
||||
|
||||
def get_avg_lr(optimizer):
|
||||
# Calculate the weighted average effective learning rate
|
||||
total_lr = 0
|
||||
total_params = 0
|
||||
for group in optimizer.param_groups:
|
||||
d = group['d']
|
||||
lr = group['lr']
|
||||
bias_correction = 1 # Default value
|
||||
if group['use_bias_correction']:
|
||||
beta1, beta2 = group['betas']
|
||||
k = group['k']
|
||||
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
|
||||
|
||||
effective_lr = d * lr * bias_correction
|
||||
|
||||
# Count the number of parameters in this group
|
||||
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
|
||||
total_lr += effective_lr * num_params
|
||||
total_params += num_params
|
||||
|
||||
if total_params == 0:
|
||||
return 0.0
|
||||
else: return total_lr / total_params
|
||||
|
||||
|
||||
import re
|
||||
|
||||
def replace_in_string(s, replacements):
|
||||
while True:
|
||||
replaced = False
|
||||
for target, replacement in replacements.items():
|
||||
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
|
||||
if new_s != s:
|
||||
s = new_s
|
||||
replaced = True
|
||||
if not replaced:
|
||||
break
|
||||
return s
|
||||
|
||||
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
|
||||
if "_no_token" in lora_path:
|
||||
return prompt
|
||||
|
||||
orig_prompt = prompt
|
||||
|
||||
# Helper function to read JSON
|
||||
def read_json_from_path(path):
|
||||
with open(path, "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
# Check existence of "special_params.json"
|
||||
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
|
||||
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
|
||||
|
||||
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
|
||||
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
|
||||
|
||||
try:
|
||||
lora_name = str(training_args["name"])
|
||||
except: # fallback for old loras that dont have the name field:
|
||||
return training_args["trigger_text"] + ", " + prompt
|
||||
|
||||
lora_name_encapsulated = "<" + lora_name + ">"
|
||||
trigger_text = training_args["trigger_text"]
|
||||
|
||||
try:
|
||||
mode = training_args["concept_mode"]
|
||||
except KeyError:
|
||||
try:
|
||||
mode = training_args["mode"]
|
||||
except KeyError:
|
||||
mode = "object"
|
||||
|
||||
# Handle different modes
|
||||
if mode != "style":
|
||||
replacements = {
|
||||
"<concept>": trigger_text,
|
||||
"<concepts>": trigger_text + "'s",
|
||||
lora_name_encapsulated: trigger_text,
|
||||
lora_name_encapsulated.lower(): trigger_text,
|
||||
lora_name: trigger_text,
|
||||
lora_name.lower(): trigger_text,
|
||||
}
|
||||
prompt = replace_in_string(prompt, replacements)
|
||||
if trigger_text not in prompt:
|
||||
prompt = trigger_text + ", " + prompt
|
||||
else:
|
||||
style_replacements = {
|
||||
"in the style of <concept>": "in the style of TOK",
|
||||
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
|
||||
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
|
||||
f"in the style of {lora_name}": "in the style of TOK",
|
||||
f"in the style of {lora_name.lower()}": "in the style of TOK"
|
||||
}
|
||||
prompt = replace_in_string(prompt, style_replacements)
|
||||
if "in the style of TOK" not in prompt:
|
||||
prompt = "in the style of TOK, " + prompt
|
||||
|
||||
# Final cleanup
|
||||
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
|
||||
|
||||
if interpolation and mode != "style":
|
||||
prompt = "TOK, " + prompt
|
||||
|
||||
# Replace tokens based on token map
|
||||
prompt = replace_in_string(prompt, token_map)
|
||||
|
||||
# Fix common mistakes
|
||||
fix_replacements = {
|
||||
r",,": ",",
|
||||
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
|
||||
r"\s\.": ".",
|
||||
r"\s,": ","
|
||||
}
|
||||
prompt = replace_in_string(prompt, fix_replacements)
|
||||
|
||||
if verbose:
|
||||
print('-------------------------')
|
||||
print("Adjusted prompt for LORA:")
|
||||
print(orig_prompt)
|
||||
print('-- to:')
|
||||
print(prompt)
|
||||
print('-------------------------')
|
||||
|
||||
return prompt
|
||||
|
||||
from val_prompts import val_prompts
|
||||
@torch.no_grad()
|
||||
def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0"):
|
||||
|
||||
random.seed(seed)
|
||||
|
||||
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
|
||||
training_args = json.load(f)
|
||||
concept_mode = training_args["concept_mode"]
|
||||
|
||||
if concept_mode == "style":
|
||||
validation_prompts_raw = random.sample(val_prompts['style'], n_imgs)
|
||||
validation_prompts_raw[0] = ''
|
||||
|
||||
elif concept_mode == "face":
|
||||
validation_prompts_raw = random.sample(val_prompts['face'], n_imgs)
|
||||
validation_prompts_raw[0] = '<concept>'
|
||||
else:
|
||||
validation_prompts_raw = random.sample(val_prompts['object'], n_imgs)
|
||||
validation_prompts_raw[0] = '<concept>'
|
||||
|
||||
|
||||
reload_entire_pipeline = False
|
||||
if reload_entire_pipeline: # reload the entire pipeline from disk and load in the lora module
|
||||
print(f"Reloading entire pipeline from disk..")
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
(pipeline,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
noise_scheduler,
|
||||
text_encoder_one,
|
||||
text_encoder_two,
|
||||
vae,
|
||||
unet) = load_models(pretrained_model, device, torch.float16)
|
||||
|
||||
pipeline = pipeline.to(device)
|
||||
pipeline = patch_pipe_with_lora(pipeline, lora_path)
|
||||
|
||||
else:
|
||||
print(f"Re-using training pipeline for inference, just swapping the scheduler..")
|
||||
pipeline = training_pipeline
|
||||
training_scheduler = pipeline.scheduler
|
||||
|
||||
pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config)
|
||||
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path) for prompt in validation_prompts_raw]
|
||||
generator = torch.Generator(device=device).manual_seed(0)
|
||||
pipeline_args = {
|
||||
"negative_prompt": "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft",
|
||||
"num_inference_steps": n_steps,
|
||||
"guidance_scale": 7,
|
||||
"height": render_size[0],
|
||||
"width": render_size[1],
|
||||
}
|
||||
|
||||
if is_lora > 0:
|
||||
cross_attention_kwargs = {"scale": lora_scale}
|
||||
else:
|
||||
cross_attention_kwargs = None
|
||||
|
||||
for i in range(n_imgs):
|
||||
pipeline_args["prompt"] = validation_prompts[i]
|
||||
print(f"Rendering validation img with prompt: {validation_prompts[i]}")
|
||||
image = pipeline(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
|
||||
image.save(os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"), format="JPEG", quality=95)
|
||||
|
||||
# create img_grid:
|
||||
img_grid_path = make_validation_img_grid(lora_path)
|
||||
|
||||
if not reload_entire_pipeline: # restore the training scheduler
|
||||
pipeline.scheduler = training_scheduler
|
||||
|
||||
return validation_prompts_raw
|
||||
|
||||
|
||||
def main(
|
||||
pretrained_model,
|
||||
instance_data_dir: Optional[str] = "./dataset/zeke/captions.csv",
|
||||
output_dir: str = "lora_output",
|
||||
seed: Optional[int] = random.randint(0, 2**32 - 1),
|
||||
resolution: int = 768,
|
||||
crops_coords_top_left_h: int = 0,
|
||||
crops_coords_top_left_w: int = 0,
|
||||
train_batch_size: int = 1,
|
||||
do_cache: bool = True,
|
||||
num_train_epochs: int = 10000,
|
||||
max_train_steps: Optional[int] = None,
|
||||
checkpointing_steps: int = 500000, # default to no checkpoints
|
||||
gradient_accumulation_steps: int = 1, # todo
|
||||
unet_learning_rate: float = 1.0,
|
||||
ti_lr: float = 3e-4,
|
||||
lora_lr: float = 1.0,
|
||||
prodigy_d_coef: float = 0.33,
|
||||
l1_penalty: float = 0.0,
|
||||
lora_weight_decay: float = 0.005,
|
||||
ti_weight_decay: float = 0.001,
|
||||
scale_lr: bool = False,
|
||||
lr_scheduler: str = "constant",
|
||||
lr_warmup_steps: int = 50,
|
||||
lr_num_cycles: int = 1,
|
||||
lr_power: float = 1.0,
|
||||
snr_gamma: float = 5.0,
|
||||
dataloader_num_workers: int = 0,
|
||||
allow_tf32: bool = True,
|
||||
mixed_precision: Optional[str] = "bf16",
|
||||
device: str = "cuda:0",
|
||||
token_dict: dict = {"TOKEN": "<s0>"},
|
||||
inserting_list_tokens: List[str] = ["<s0>"],
|
||||
verbose: bool = True,
|
||||
is_lora: bool = True,
|
||||
lora_rank: int = 8,
|
||||
args_dict: dict = {},
|
||||
debug: bool = False,
|
||||
hard_pivot: bool = True,
|
||||
off_ratio_power: float = 0.1,
|
||||
) -> None:
|
||||
if allow_tf32:
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
print("Using seed", seed)
|
||||
torch.manual_seed(seed)
|
||||
|
||||
weight_dtype = torch.float32
|
||||
if mixed_precision == "fp16":
|
||||
weight_dtype = torch.float16
|
||||
elif mixed_precision == "bf16":
|
||||
weight_dtype = torch.bfloat16
|
||||
|
||||
print(f"Loading models with weight_dtype: {weight_dtype}")
|
||||
|
||||
if scale_lr:
|
||||
unet_learning_rate = (
|
||||
unet_learning_rate * gradient_accumulation_steps * train_batch_size
|
||||
)
|
||||
|
||||
(
|
||||
pipe,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
noise_scheduler,
|
||||
text_encoder_one,
|
||||
text_encoder_two,
|
||||
vae,
|
||||
unet,
|
||||
) = load_models(pretrained_model, device, weight_dtype)
|
||||
|
||||
# Initialize new tokens for training.
|
||||
embedding_handler = TokenEmbeddingsHandler(
|
||||
[text_encoder_one, text_encoder_two], [tokenizer_one, tokenizer_two]
|
||||
)
|
||||
|
||||
#starting_toks = ["person", "face"]
|
||||
starting_toks = None
|
||||
embedding_handler.initialize_new_tokens(inserting_toks=inserting_list_tokens, starting_toks=starting_toks, seed=seed)
|
||||
text_encoders = [text_encoder_one, text_encoder_two]
|
||||
|
||||
unet_param_to_optimize = []
|
||||
text_encoder_parameters = []
|
||||
for text_encoder in text_encoders:
|
||||
if text_encoder is not None:
|
||||
for name, param in text_encoder.named_parameters():
|
||||
if "token_embedding" in name:
|
||||
param.requires_grad = True
|
||||
text_encoder_parameters.append(param)
|
||||
else:
|
||||
param.requires_grad = False
|
||||
|
||||
unet_param_to_optimize_names = []
|
||||
unet_lora_parameters = []
|
||||
|
||||
if not is_lora:
|
||||
WHITELIST_PATTERNS = [
|
||||
# "*.attn*.weight",
|
||||
# "*ff*.weight",
|
||||
"*"
|
||||
]
|
||||
BLACKLIST_PATTERNS = ["*.norm*.weight", "*time*"]
|
||||
for name, param in unet.named_parameters():
|
||||
if any(
|
||||
fnmatch.fnmatch(name, pattern) for pattern in WHITELIST_PATTERNS
|
||||
) and not any(
|
||||
fnmatch.fnmatch(name, pattern) for pattern in BLACKLIST_PATTERNS
|
||||
):
|
||||
param.requires_grad_(True)
|
||||
unet_param_to_optimize_names.append(name)
|
||||
print(f"Training: {name}")
|
||||
else:
|
||||
param.requires_grad_(False)
|
||||
|
||||
# Optimizer creation
|
||||
params_to_optimize = [
|
||||
{
|
||||
"params": text_encoder_parameters,
|
||||
"lr": ti_lr,
|
||||
"weight_decay": ti_weight_decay,
|
||||
},
|
||||
]
|
||||
|
||||
params_to_optimize_prodigy = [
|
||||
{
|
||||
"params": unet_param_to_optimize,
|
||||
"lr": unet_learning_rate,
|
||||
"weight_decay": lora_weight_decay,
|
||||
},
|
||||
]
|
||||
|
||||
else:
|
||||
|
||||
# Do lora-training instead.
|
||||
unet.requires_grad_(False)
|
||||
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
|
||||
unet_lora_config = LoraConfig(
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_rank,
|
||||
init_lora_weights="gaussian",
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
#use_rslora=True,
|
||||
use_dora=True,
|
||||
)
|
||||
#unet.add_adapter(unet_lora_config)
|
||||
|
||||
unet = get_peft_model(unet, unet_lora_config)
|
||||
print_trainable_parameters(unet, name = 'unet')
|
||||
|
||||
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
|
||||
|
||||
params_to_optimize = [
|
||||
{
|
||||
"params": text_encoder_parameters,
|
||||
"lr": ti_lr,
|
||||
"weight_decay": ti_weight_decay,
|
||||
},
|
||||
]
|
||||
|
||||
params_to_optimize_prodigy = [
|
||||
{
|
||||
"params": unet_lora_parameters,
|
||||
"lr": 1.0,
|
||||
"weight_decay": lora_weight_decay,
|
||||
},
|
||||
]
|
||||
|
||||
optimizer_type = "prodigy" # hardcode for now
|
||||
|
||||
if optimizer_type != "prodigy":
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
weight_decay=0.0, # this wd doesn't matter, I think
|
||||
)
|
||||
else:
|
||||
try:
|
||||
import prodigyopt
|
||||
except ImportError:
|
||||
raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`")
|
||||
|
||||
# Note: the specific settings of Prodigy seem to matter A LOT
|
||||
optimizer_prod = prodigyopt.Prodigy(
|
||||
params_to_optimize_prodigy,
|
||||
d_coef = prodigy_d_coef,
|
||||
lr=1.0,
|
||||
decouple=True,
|
||||
use_bias_correction=True,
|
||||
safeguard_warmup=True,
|
||||
weight_decay=lora_weight_decay,
|
||||
betas=(0.9, 0.99),
|
||||
growth_rate=1.025, # this slows down the lr_rampup
|
||||
)
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
weight_decay=ti_weight_decay,
|
||||
)
|
||||
|
||||
train_dataset = PreprocessedDataset(
|
||||
instance_data_dir,
|
||||
tokenizer_one,
|
||||
tokenizer_two,
|
||||
vae,
|
||||
do_cache=True,
|
||||
substitute_caption_map=token_dict,
|
||||
)
|
||||
|
||||
print(f"# PTI : Loaded dataset, do_cache: {do_cache}")
|
||||
train_dataloader = torch.utils.data.DataLoader(
|
||||
train_dataset,
|
||||
batch_size=train_batch_size,
|
||||
shuffle=True,
|
||||
num_workers=dataloader_num_workers,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(
|
||||
len(train_dataloader) / gradient_accumulation_steps
|
||||
)
|
||||
if max_train_steps is None:
|
||||
max_train_steps = num_train_epochs * num_update_steps_per_epoch
|
||||
|
||||
lr_scheduler = get_scheduler(
|
||||
lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=lr_warmup_steps * gradient_accumulation_steps,
|
||||
num_training_steps=max_train_steps * gradient_accumulation_steps,
|
||||
num_cycles=lr_num_cycles,
|
||||
power=lr_power,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(
|
||||
len(train_dataloader) / gradient_accumulation_steps
|
||||
)
|
||||
num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
total_batch_size = train_batch_size * gradient_accumulation_steps
|
||||
|
||||
if verbose:
|
||||
print(f"# PTI : Running training ")
|
||||
print(f"# PTI : Num examples = {len(train_dataset)}")
|
||||
print(f"# PTI : Num batches each epoch = {len(train_dataloader)}")
|
||||
print(f"# PTI : Num Epochs = {num_train_epochs}")
|
||||
print(f"# PTI : Instantaneous batch size per device = {train_batch_size}")
|
||||
print(
|
||||
f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}"
|
||||
)
|
||||
print(f"# PTI : Gradient Accumulation steps = {gradient_accumulation_steps}")
|
||||
print(f"# PTI : Total optimization steps = {max_train_steps}")
|
||||
|
||||
global_step = 0
|
||||
first_epoch = 0
|
||||
last_save_step = 0
|
||||
|
||||
progress_bar = tqdm(range(global_step, max_train_steps), position=0, leave=True)
|
||||
checkpoint_dir = os.path.join(str(output_dir), "checkpoints")
|
||||
if os.path.exists(checkpoint_dir):
|
||||
shutil.rmtree(checkpoint_dir)
|
||||
os.makedirs(f"{checkpoint_dir}")
|
||||
|
||||
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
|
||||
#embedding_handler.pre_optimize_token_embeddings(train_dataset)
|
||||
|
||||
ti_lrs, lora_lrs = [], []
|
||||
losses = []
|
||||
start_time, images_done = time.time(), 0
|
||||
|
||||
for epoch in range(first_epoch, num_train_epochs):
|
||||
unet.train()
|
||||
progress_bar.set_description(f"# PTI :step: {global_step}, epoch: {epoch}")
|
||||
|
||||
for step, batch in enumerate(train_dataloader):
|
||||
progress_bar.update(1)
|
||||
|
||||
if hard_pivot:
|
||||
if epoch >= num_train_epochs // 2:
|
||||
if optimizer is not None:
|
||||
print("----------------------")
|
||||
print("# PTI : Pivot halfway")
|
||||
print("----------------------")
|
||||
# remove text encoder parameters from the optimizer
|
||||
optimizer.param_groups = None
|
||||
# remove the optimizer state corresponding to text_encoder_parameters
|
||||
for param in text_encoder_parameters:
|
||||
if param in optimizer.state:
|
||||
del optimizer.state[param]
|
||||
optimizer = None
|
||||
|
||||
else: # Update learning rates gradually:
|
||||
finegrained_epoch = epoch + step / len(train_dataloader)
|
||||
completion_f = finegrained_epoch / num_train_epochs
|
||||
# param_groups[1] goes from ti_lr to 0.0 over the course of training
|
||||
optimizer.param_groups[0]['lr'] = ti_lr * (1 - completion_f) ** 2.0
|
||||
|
||||
|
||||
try: #sdxl
|
||||
(tok1, tok2), vae_latent, mask = batch
|
||||
except: #sd15
|
||||
tok1, vae_latent, mask = batch
|
||||
tok2 = None
|
||||
|
||||
vae_latent = vae_latent.to(weight_dtype)
|
||||
|
||||
# tokens to text embeds
|
||||
prompt_embeds_list = []
|
||||
for tok, text_encoder in zip((tok1, tok2), text_encoders):
|
||||
if tok is None:
|
||||
continue
|
||||
|
||||
prompt_embeds_out = text_encoder(
|
||||
tok.to(text_encoder.device),
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
pooled_prompt_embeds = prompt_embeds_out[0]
|
||||
prompt_embeds = prompt_embeds_out.hidden_states[-2]
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
|
||||
prompt_embeds_list.append(prompt_embeds)
|
||||
|
||||
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
|
||||
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
|
||||
|
||||
# Create Spatial-dimensional conditions.
|
||||
original_size = (resolution, resolution)
|
||||
target_size = (resolution, resolution)
|
||||
crops_coords_top_left = (crops_coords_top_left_h, crops_coords_top_left_w)
|
||||
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(device, dtype=prompt_embeds.dtype).repeat(
|
||||
bs_embed, 1
|
||||
)
|
||||
|
||||
# Sample noise that we'll add to the latents:
|
||||
noise = torch.randn_like(vae_latent)
|
||||
|
||||
noise_offset = 0.05 # TODO, turn this into an input arg and do a grid search
|
||||
if noise_offset > 0.0:
|
||||
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
|
||||
noise += noise_offset * torch.randn(
|
||||
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
|
||||
|
||||
bsz = vae_latent.shape[0]
|
||||
|
||||
timesteps = torch.randint(
|
||||
0,
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
(bsz,),
|
||||
device=vae_latent.device,
|
||||
).long()
|
||||
|
||||
noisy_model_input = noise_scheduler.add_noise(vae_latent, noise, timesteps)
|
||||
|
||||
noise_sigma = 0.0
|
||||
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
|
||||
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,1:-2,:]) * noise_sigma
|
||||
|
||||
# Predict the noise residual
|
||||
model_pred = unet(
|
||||
noisy_model_input,
|
||||
timesteps,
|
||||
prompt_embeds,
|
||||
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
|
||||
).sample
|
||||
|
||||
# Get the unet prediction target depending on the prediction type:
|
||||
if noise_scheduler.config.prediction_type == "epsilon":
|
||||
target = noise
|
||||
elif noise_scheduler.config.prediction_type == "v_prediction":
|
||||
target = noise_scheduler.get_velocity(model_input, noise, timesteps)
|
||||
else:
|
||||
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
|
||||
|
||||
# Compute the loss:
|
||||
if snr_gamma is None:
|
||||
loss = (model_pred - target).pow(2) * mask
|
||||
|
||||
# modulate loss by the inverse of the mask's mean value
|
||||
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
|
||||
mean_mask_values = mean_mask_values / mean_mask_values.mean()
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
|
||||
|
||||
# Average the normalized errors across the batch
|
||||
loss = loss.mean()
|
||||
|
||||
else:
|
||||
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
|
||||
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
|
||||
# This is discussed in Section 4.2 of the same paper.
|
||||
snr = compute_snr(noise_scheduler, timesteps)
|
||||
base_weight = (
|
||||
torch.stack([snr, snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
|
||||
)
|
||||
if noise_scheduler.config.prediction_type == "v_prediction":
|
||||
# Velocity objective needs to be floored to an SNR weight of one.
|
||||
mse_loss_weights = base_weight + 1
|
||||
else:
|
||||
# Epsilon and sample both use the same loss weights.
|
||||
mse_loss_weights = base_weight
|
||||
|
||||
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
|
||||
loss = (model_pred - target).pow(2) * mask
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
|
||||
|
||||
if 1: # modulate loss by the inverse of the mask's mean value
|
||||
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
|
||||
mean_mask_values = mean_mask_values / mean_mask_values.mean()
|
||||
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
|
||||
|
||||
loss = loss.mean()
|
||||
|
||||
if l1_penalty > 0.0:
|
||||
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
|
||||
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
|
||||
loss += l1_penalty * l1_norm
|
||||
|
||||
losses.append(loss.item())
|
||||
|
||||
loss = loss / gradient_accumulation_steps
|
||||
loss.backward()
|
||||
|
||||
'''
|
||||
apart from the usual gradient accumulation steps,
|
||||
we also do a backward pass after computing the last forward pass in the epoch (last_batch == True)
|
||||
this is to make sure that we're not missing out on any data
|
||||
'''
|
||||
last_batch = (step + 1 == len(train_dataloader))
|
||||
if (step + 1) % gradient_accumulation_steps == 0 or last_batch:
|
||||
if optimizer is not None:
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
optimizer_prod.step()
|
||||
optimizer_prod.zero_grad()
|
||||
|
||||
# after every optimizer step, we reset the non-trainable embeddings to the original embeddings
|
||||
embedding_handler.retract_embeddings(print_stds = (global_step % 50 == 0))
|
||||
embedding_handler.fix_embedding_std(off_ratio_power)
|
||||
|
||||
# Track the learning rates for final plotting:
|
||||
lora_lrs.append(get_avg_lr(optimizer_prod))
|
||||
try:
|
||||
ti_lrs.append(optimizer.param_groups[0]['lr'])
|
||||
except:
|
||||
ti_lrs.append(0.0)
|
||||
|
||||
# Print some statistics:
|
||||
if (global_step % checkpointing_steps == 0):
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
||||
save_lora(output_save_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names)
|
||||
last_save_step = global_step
|
||||
|
||||
if debug:
|
||||
token_embeddings = embedding_handler.get_trainable_embeddings()
|
||||
for i, token_embeddings_i in enumerate(token_embeddings):
|
||||
plot_torch_hist(token_embeddings_i[0], global_step, output_dir, f"embeddings_weights_token_0_{i}", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
|
||||
plot_torch_hist(token_embeddings_i[1], global_step, output_dir, f"embeddings_weights_token_1_{i}", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
|
||||
|
||||
embedding_handler.print_token_info()
|
||||
plot_torch_hist(unet_lora_parameters, global_step, output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
|
||||
plot_loss(losses, save_path=f'{output_dir}/losses.png')
|
||||
plot_lrs(lora_lrs, ti_lrs, save_path=f'{output_dir}/learning_rates.png')
|
||||
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, seed, is_lora, pretrained_model, n_imgs = 4)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
images_done += train_batch_size
|
||||
global_step += 1
|
||||
|
||||
if global_step % 100 == 0:
|
||||
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
|
||||
|
||||
if global_step % (max_train_steps//20) == 0:
|
||||
progress = (global_step / max_train_steps) + 0.05
|
||||
yield np.min((progress, 1.0))
|
||||
|
||||
|
||||
# final_save
|
||||
if (global_step - last_save_step) > 51:
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
||||
else:
|
||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
|
||||
|
||||
if debug:
|
||||
plot_loss(losses, save_path=f'{output_dir}/losses.png')
|
||||
plot_lrs(lora_lrs, ti_lrs, save_path=f'{output_dir}/learning_rates.png')
|
||||
plot_torch_hist(unet_lora_parameters, global_step, output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
|
||||
plot_torch_hist(embedding_handler.get_trainable_embeddings(), global_step, output_dir, "embeddings_weights", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
|
||||
|
||||
if not os.path.exists(output_save_dir):
|
||||
save_lora(output_save_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names)
|
||||
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, seed, is_lora, pretrained_model, n_imgs = 4, n_steps = 35)
|
||||
else:
|
||||
print(f"Skipping final save, {output_save_dir} already exists")
|
||||
|
||||
del unet
|
||||
del vae
|
||||
del text_encoder_one
|
||||
del text_encoder_two
|
||||
del tokenizer_one
|
||||
del tokenizer_two
|
||||
del embedding_handler
|
||||
del pipe
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
with open(f"{output_save_dir}/training_args.json", "w") as f:
|
||||
args_dict["grid_prompts"] = validation_prompts
|
||||
json.dump(args_dict, f, indent=4)
|
||||
|
||||
return output_save_dir, validation_prompts
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
-27
@@ -1,27 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Check if the target directory is provided
|
||||
if [ -z "$1" ]; then
|
||||
echo "Usage: $0 target_directory"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Check if the target directory exists
|
||||
if [ ! -d "$1" ]; then
|
||||
echo "Error: directory '$1' doesn't exist."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Navigate to the target directory
|
||||
cd "$1" || exit 1
|
||||
|
||||
# Initialize empty zip file
|
||||
zip -r9 "args.zip" --exclude=*
|
||||
|
||||
# Find and zip all training_args.json files
|
||||
find . -type f -name 'training_args.json' -exec zip -r "args.zip" {} +
|
||||
|
||||
# Navigate back to the original directory
|
||||
cd - || exit 1
|
||||
|
||||
echo "All training_args.json files have been zipped into args.zip"
|
||||
Reference in New Issue
Block a user