Compare commits
185 Commits
developmen
...
4fa8fac5fd
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4fa8fac5fd | ||
|
|
a48c9aba8d | ||
|
|
60232def91 | ||
|
|
3843e0d148 | ||
|
|
e127c079da | ||
|
|
34db804c76 | ||
|
|
4d35a29c97 | ||
|
|
b322d05fa3 | ||
|
|
8577849eeb | ||
|
|
338c77d677 | ||
|
|
e07a98a50c | ||
|
|
6a754b2710 | ||
|
|
a939cf3730 | ||
|
|
169dbd22ba | ||
|
|
6e7d721382 | ||
|
|
dc6f36cd82 | ||
|
|
5603f9e004 | ||
|
|
c45887192a | ||
|
|
13a965a26c | ||
|
|
77ee7090e8 | ||
|
|
078396ceac | ||
|
|
f944eeaa4d | ||
|
|
81899310f8 | ||
|
|
f9179540d2 | ||
|
|
452e0e286d | ||
|
|
165510ace2 | ||
|
|
0355662e8e | ||
|
|
b99d36dfdb | ||
|
|
9001e5c933 | ||
|
|
7fed4ea761 | ||
|
|
e07bf11727 | ||
|
|
c728cc9a0b | ||
|
|
00bd3d54a3 | ||
|
|
f7cf2f866f | ||
|
|
465bc1e2f8 | ||
|
|
0beca0d4a7 | ||
|
|
418f5f7e8c | ||
|
|
9ee1ef2a0a | ||
|
|
599fafe01f | ||
|
|
af108bb964 | ||
|
|
89d61a3b8e | ||
|
|
a6aa4b2c7d | ||
|
|
f8f0657b68 | ||
|
|
7f0ecdb377 | ||
|
|
fbed8568fb | ||
|
|
6d31c6db73 | ||
|
|
6490a326e5 | ||
|
|
8d48ad4e85 | ||
|
|
ec1ea7aa0e | ||
|
|
fa02e774b0 | ||
|
|
2308ef2868 | ||
|
|
b3e03295ad | ||
|
|
e69a520616 | ||
|
|
acafe9984f | ||
|
|
653fe60f16 | ||
|
|
c2424087d6 | ||
|
|
272c8608c2 | ||
|
|
99f24cfb0c | ||
|
|
187663ab55 | ||
|
|
edb7e827ee | ||
|
|
0ea27011d5 | ||
|
|
f321de7bdb | ||
|
|
88acc28d7f | ||
|
|
de2da96a81 | ||
|
|
9beea1c268 | ||
|
|
369aa143bc | ||
|
|
87ba867fdc | ||
|
|
03613c523f | ||
|
|
47744373f2 | ||
|
|
443c996e7f | ||
|
|
8f0f467c20 | ||
|
|
e81e19fd0f | ||
|
|
0bc4d555c7 | ||
|
|
80aa2dbb80 | ||
|
|
8d799031cf | ||
|
|
6e92922c14 | ||
|
|
c51235c486 | ||
|
|
c2c4e8cf34 | ||
|
|
4c249cf607 | ||
|
|
c2d5f712a3 | ||
|
|
22d2f6e28f | ||
|
|
a2301cf28c | ||
|
|
11e426fdf1 | ||
|
|
58dffd43a8 | ||
|
|
e4558dff4b | ||
|
|
c062b7716c | ||
|
|
c008405480 | ||
|
|
93e5df1d59 | ||
|
|
045e4a6e15 | ||
|
|
76f225a467 | ||
|
|
cab8a1c7b8 | ||
|
|
acb06d6ff3 | ||
|
|
bb57623a35 | ||
|
|
3072d20f17 | ||
|
|
f6b21f47bb | ||
|
|
603ceca3ca | ||
|
|
657fd09f25 | ||
|
|
8407c4deea | ||
|
|
64f2b085b7 | ||
|
|
7165f2d25a | ||
|
|
5d47244c57 | ||
|
|
ada722c9e4 | ||
|
|
696f73c30d | ||
|
|
e3410413b9 | ||
|
|
37cebd9458 | ||
|
|
bd10d2d668 | ||
|
|
cb5d28cba9 | ||
|
|
3f3636b788 | ||
|
|
833c833f28 | ||
|
|
68b7e159bc | ||
|
|
5a45c709cd | ||
|
|
10e1ecf1e8 | ||
|
|
b96913d73c | ||
|
|
5da3613e0b | ||
|
|
5a70b7f38d | ||
|
|
377b81ee3e | ||
|
|
2d0a1be59d | ||
|
|
7284aab7c0 | ||
|
|
427847ac4c | ||
|
|
9c1cc9641e | ||
|
|
89f4bcad2e | ||
|
|
016687bda1 | ||
|
|
72de68d8aa | ||
|
|
d87b49882c | ||
|
|
f415bac7b5 | ||
|
|
f1cb87fe9e | ||
|
|
8f9cd823d1 | ||
|
|
b01e8d889a | ||
|
|
1325613583 | ||
|
|
337945de9a | ||
|
|
561914d8e6 | ||
|
|
b0a0f28191 | ||
|
|
f965a1299f | ||
|
|
1bd94f0f01 | ||
|
|
9ffa8c3711 | ||
|
|
b68c3ef734 | ||
|
|
49c41e6a5f | ||
|
|
2478554c95 | ||
|
|
93b52932c1 | ||
|
|
4ec4025cbb | ||
|
|
e074058faa | ||
|
|
a8481c1670 | ||
|
|
e18e0cb5f8 | ||
|
|
177c7130ec | ||
|
|
1ae1017748 | ||
|
|
92b9c71d44 | ||
|
|
f17ad8d794 | ||
|
|
86c70a2a1f | ||
|
|
655533d4c7 | ||
|
|
eebd3c8212 | ||
|
|
5276975fb0 | ||
|
|
e190fbaeb8 | ||
|
|
290393f7ae | ||
|
|
b2a54c8f36 | ||
|
|
b767d29b3c | ||
|
|
645b27f97a | ||
|
|
65c08b09c3 | ||
|
|
afc231efc1 | ||
|
|
bafacf3b65 | ||
|
|
0892dec4a5 | ||
|
|
eeee4a1620 | ||
|
|
d11ed7f66c | ||
|
|
27ad79053e | ||
|
|
05ae95ca89 | ||
|
|
0f8daa5612 | ||
|
|
7703e3a15e | ||
|
|
0f597f453e | ||
|
|
dfb64b5957 | ||
|
|
82098e5d6e | ||
|
|
b653906715 | ||
|
|
13d32423f6 | ||
|
|
39870411d8 | ||
|
|
e5177833b2 | ||
|
|
eaa0fb6253 | ||
|
|
eaec2f5a52 | ||
|
|
92cb5ae096 | ||
|
|
bd2bce9b92 | ||
|
|
537af79b0d | ||
|
|
7624241032 | ||
|
|
0d5943af91 | ||
|
|
be815f9c47 | ||
|
|
3443d6aafa | ||
|
|
bef10a639c | ||
|
|
3eb3535683 | ||
|
|
3d387103cd |
20
.github/ISSUE_TEMPLATE/bug_report.md
vendored
Normal file
20
.github/ISSUE_TEMPLATE/bug_report.md
vendored
Normal file
@@ -0,0 +1,20 @@
|
||||
---
|
||||
name: Bug Report
|
||||
about: For bugs only. Not for feature requests or questions.
|
||||
title: ''
|
||||
labels: ''
|
||||
assignees: ''
|
||||
---
|
||||
|
||||
## This is for bugs only
|
||||
|
||||
Did you already ask [in the discord](https://discord.gg/VXmU2f5WEU)?
|
||||
|
||||
Yes/No
|
||||
|
||||
You verified that this is a bug and not a feature request or question by asking [in the discord](https://discord.gg/VXmU2f5WEU)?
|
||||
|
||||
Yes/No
|
||||
|
||||
## Describe the bug
|
||||
|
||||
5
.github/ISSUE_TEMPLATE/config.yml
vendored
Normal file
5
.github/ISSUE_TEMPLATE/config.yml
vendored
Normal file
@@ -0,0 +1,5 @@
|
||||
blank_issues_enabled: false
|
||||
contact_links:
|
||||
- name: Ask in the Discord BEFORE opening an issue
|
||||
url: https://discord.gg/VXmU2f5WEU
|
||||
about: Please ask in the discord before opening a github issue.
|
||||
3
.gitignore
vendored
3
.gitignore
vendored
@@ -172,4 +172,5 @@ cython_debug/
|
||||
/output/*
|
||||
!/output/.gitkeep
|
||||
/extensions/*
|
||||
!/extensions/example
|
||||
!/extensions/example
|
||||
/temp
|
||||
10
FAQ.md
Normal file
10
FAQ.md
Normal file
@@ -0,0 +1,10 @@
|
||||
# FAQ
|
||||
|
||||
WIP. Will continue to add things as they are needed.
|
||||
|
||||
## FLUX.1 Training
|
||||
|
||||
#### How much VRAM is required to train a lora on FLUX.1?
|
||||
|
||||
24GB minimum is required.
|
||||
|
||||
21
LICENSE
Normal file
21
LICENSE
Normal file
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 Ostris, LLC
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
202
README.md
202
README.md
@@ -1,16 +1,21 @@
|
||||
# AI Toolkit by Ostris
|
||||
|
||||
## IMPORTANT NOTE - READ THIS
|
||||
This is an active WIP repo that is not ready for others to use. And definitely not ready for non developers to use.
|
||||
I am making major breaking changes and pushing straight to master until I have it in a planned state. I have big changes
|
||||
planned for config files and the general structure. I may change how training works entirely. You are welcome to use
|
||||
but keep that in mind. If more people start to use it, I will follow better branch checkout standards, but for now
|
||||
this is my personal active experiment.
|
||||
This is my research repo. I do a lot of experiments in it and it is possible that I will break things.
|
||||
If something breaks, checkout an earlier commit. This repo can train a lot of things, and it is
|
||||
hard to keep up with all of them.
|
||||
|
||||
Report bugs as you find them, but not knowing how to train ML models, setup an environment, or use python is not a bug.
|
||||
I will make all of this more user-friendly eventually
|
||||
## Support my work
|
||||
|
||||
I will make a better readme later.
|
||||
<a href="https://glif.app" target="_blank">
|
||||
<img alt="glif.app" src="https://raw.githubusercontent.com/ostris/ai-toolkit/main/assets/glif.svg?v=1" width="256" height="auto">
|
||||
</a>
|
||||
|
||||
|
||||
My work on this project would not be possible without the amazing support of [Glif](https://glif.app/) and everyone on the
|
||||
team. If you want to support me, support Glif. [Join the site](https://glif.app/),
|
||||
[Join us on Discord](https://discord.com/invite/nuR9zZ2nsh), [follow us on Twitter](https://x.com/heyglif)
|
||||
and come make some cool stuff with us
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -30,8 +35,8 @@ git submodule update --init --recursive
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
# .\venv\Scripts\activate on windows
|
||||
# windows install pytorch first with
|
||||
# pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu117
|
||||
# install torch first
|
||||
pip3 install torch
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
@@ -42,17 +47,186 @@ cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
.\venv\Scripts\activate
|
||||
pip install torch --use-pep517 --extra-index-url https://download.pytorch.org/whl/cu118
|
||||
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
## FLUX.1 Training
|
||||
|
||||
### Tutorial
|
||||
|
||||
To get started quickly, check out [@araminta_k](https://x.com/araminta_k) tutorial on [Finetuning Flux Dev on a 3090](https://www.youtube.com/watch?v=HzGW_Kyermg) with 24GB VRAM.
|
||||
|
||||
|
||||
### Requirements
|
||||
You currently need a GPU with **at least 24GB of VRAM** to train FLUX.1. If you are using it as your GPU to control
|
||||
your monitors, you probably need to set the flag `low_vram: true` in the config file under `model:`. This will quantize
|
||||
the model on CPU and should allow it to train with monitors attached. Users have gotten it to work on Windows with WSL,
|
||||
but there are some reports of a bug when running on windows natively.
|
||||
I have only tested on linux for now. This is still extremely experimental
|
||||
and a lot of quantizing and tricks had to happen to get it to fit on 24GB at all.
|
||||
|
||||
### FLUX.1-dev
|
||||
|
||||
FLUX.1-dev has a non-commercial license. Which means anything you train will inherit the
|
||||
non-commercial license. It is also a gated model, so you need to accept the license on HF before using it.
|
||||
Otherwise, this will fail. Here are the required steps to setup a license.
|
||||
|
||||
1. Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)
|
||||
2. Make a file named `.env` in the root on this folder
|
||||
3. [Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and add it to the `.env` file like so `HF_TOKEN=your_key_here`
|
||||
|
||||
### FLUX.1-schnell
|
||||
|
||||
FLUX.1-schnell is Apache 2.0. Anything trained on it can be licensed however you want and it does not require a HF_TOKEN to train.
|
||||
However, it does require a special adapter to train with it, [ostris/FLUX.1-schnell-training-adapter](https://huggingface.co/ostris/FLUX.1-schnell-training-adapter).
|
||||
It is also highly experimental. For best overall quality, training on FLUX.1-dev is recommended.
|
||||
|
||||
To use it, You just need to add the assistant to the `model` section of your config file like so:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter"
|
||||
is_flux: true
|
||||
quantize: true
|
||||
```
|
||||
|
||||
You also need to adjust your sample steps since schnell does not require as many
|
||||
|
||||
```yaml
|
||||
sample:
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
```
|
||||
|
||||
### Training
|
||||
1. Copy the example config file located at `config/examples/train_lora_flux_24gb.yaml` (`config/examples/train_lora_flux_schnell_24gb.yaml` for schnell) to the `config` folder and rename it to `whatever_you_want.yml`
|
||||
2. Edit the file following the comments in the file
|
||||
3. Run the file like so `python run.py config/whatever_you_want.yml`
|
||||
|
||||
A folder with the name and the training folder from the config file will be created when you start. It will have all
|
||||
checkpoints and images in it. You can stop the training at any time using ctrl+c and when you resume, it will pick back up
|
||||
from the last checkpoint.
|
||||
|
||||
IMPORTANT. If you press crtl+c while it is saving, it will likely corrupt that checkpoint. So wait until it is done saving
|
||||
|
||||
### Need help?
|
||||
|
||||
Please do not open a bug report unless it is a bug in the code. You are welcome to [Join my Discord](https://discord.gg/VXmU2f5WEU)
|
||||
and ask for help there. However, please refrain from PMing me directly with general question or support. Ask in the discord
|
||||
and I will answer when I can.
|
||||
|
||||
## Training in RunPod
|
||||
Example RunPod template: **runpod/pytorch:2.2.0-py3.10-cuda12.1.1-devel-ubuntu22.04**
|
||||
> You need a minimum of 24GB VRAM, pick a GPU by your preference.
|
||||
|
||||
#### Example config ($0.5/hr):
|
||||
- 1x A40 (48 GB VRAM)
|
||||
- 19 vCPU 100 GB RAM
|
||||
|
||||
#### Custom overrides (you need some storage to clone FLUX.1, store datasets, store trained models and samples):
|
||||
- ~120 GB Disk
|
||||
- ~120 GB Pod Volume
|
||||
- Start Jupyter Notebook
|
||||
|
||||
### 1. Setup
|
||||
```
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
source venv/bin/activate
|
||||
pip install torch
|
||||
pip install -r requirements.txt
|
||||
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
|
||||
```
|
||||
### 2. Upload your dataset
|
||||
- Create a new folder in the root, name it `dataset` or whatever you like.
|
||||
- Drag and drop your .jpg, .jpeg, or .png images and .txt files inside the newly created dataset folder.
|
||||
|
||||
### 3. Login into Hugging Face with an Access Token
|
||||
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
|
||||
- Run ```huggingface-cli login``` and paste your token.
|
||||
|
||||
### 4. Training
|
||||
- Copy an example config file located at ```config/examples``` to the config folder and rename it to ```whatever_you_want.yml```.
|
||||
- Edit the config following the comments in the file.
|
||||
- Change ```folder_path: "/path/to/images/folder"``` to your dataset path like ```folder_path: "/workspace/ai-toolkit/your-dataset"```.
|
||||
- Run the file: ```python run.py config/whatever_you_want.yml```.
|
||||
|
||||
### Screenshot from RunPod
|
||||
<img width="1728" alt="RunPod Training Screenshot" src="https://github.com/user-attachments/assets/53a1b8ef-92fa-4481-81a7-bde45a14a7b5">
|
||||
|
||||
## Training in Modal
|
||||
|
||||
### 1. Setup
|
||||
#### ai-toolkit:
|
||||
```
|
||||
git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
git submodule update --init --recursive
|
||||
python -m venv venv
|
||||
source venv/bin/activate
|
||||
pip install torch
|
||||
pip install -r requirements.txt
|
||||
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
|
||||
```
|
||||
#### Modal:
|
||||
- Run `pip install modal` to install the modal Python package.
|
||||
- Run `modal setup` to authenticate (if this doesn’t work, try `python -m modal setup`).
|
||||
|
||||
#### Hugging Face:
|
||||
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
|
||||
- Run `huggingface-cli login` and paste your token.
|
||||
|
||||
### 2. Upload your dataset
|
||||
- Drag and drop your dataset folder containing the .jpg, .jpeg, or .png images and .txt files in `ai-toolkit`.
|
||||
|
||||
### 3. Configs
|
||||
- Copy an example config file located at ```config/examples/modal``` to the `config` folder and rename it to ```whatever_you_want.yml```.
|
||||
- Edit the config following the comments in the file, **<ins>be careful and follow the example `/root/ai-toolkit` paths</ins>**.
|
||||
|
||||
### 4. Edit run_modal.py
|
||||
- Set your entire local `ai-toolkit` path at `code_mount = modal.Mount.from_local_dir` like:
|
||||
|
||||
```
|
||||
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
|
||||
```
|
||||
- Choose a `GPU` and `Timeout` in `@app.function` _(default is A100 40GB and 2 hour timeout)_.
|
||||
|
||||
### 5. Training
|
||||
- Run the config file in your terminal: `modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml`.
|
||||
- You can monitor your training in your local terminal, or on [modal.com](https://modal.com/).
|
||||
- Models, samples and optimizer will be stored in `Storage > flux-lora-models`.
|
||||
|
||||
### 6. Saving the model
|
||||
- Check contents of the volume by running `modal volume ls flux-lora-models`.
|
||||
- Download the content by running `modal volume get flux-lora-models your-model-name`.
|
||||
- Example: `modal volume get flux-lora-models my_first_flux_lora_v1`.
|
||||
|
||||
### Screenshot from Modal
|
||||
|
||||
<img width="1728" alt="Modal Traning Screenshot" src="https://github.com/user-attachments/assets/7497eb38-0090-49d6-8ad9-9c8ea7b5388b">
|
||||
|
||||
---
|
||||
|
||||
## Current Tools
|
||||
## Dataset Preparation
|
||||
|
||||
I have so many hodge podge scripts I am going to be moving over to this that I use in my ML work. But this is what is
|
||||
here so far.
|
||||
Datasets generally need to be a folder containing images and associated text files. Currently, the only supported
|
||||
formats are jpg, jpeg, and png. Webp currently has issues. The text files should be named the same as the images
|
||||
but with a `.txt` extension. For example `image2.jpg` and `image2.txt`. The text file should contain only the caption.
|
||||
You can add the word `[trigger]` in the caption file and if you have `trigger_word` in your config, it will be automatically
|
||||
replaced.
|
||||
|
||||
Images are never upscaled but they are downscaled and placed in buckets for batching. **You do not need to crop/resize your images**.
|
||||
The loader will automatically resize them and can handle varying aspect ratios.
|
||||
|
||||
---
|
||||
|
||||
## EVERYTHING BELOW THIS LINE IS OUTDATED
|
||||
|
||||
It may still work like that, but I have not tested it in a while.
|
||||
|
||||
---
|
||||
|
||||
|
||||
40
assets/glif.svg
Normal file
40
assets/glif.svg
Normal file
@@ -0,0 +1,40 @@
|
||||
<svg width="148" height="66" viewBox="0 0 148 66" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<rect opacity="0.3" width="148" height="66" rx="33" fill="#030F2F"/>
|
||||
<g filter="url(#filter0_d_10631_12135)">
|
||||
<path d="M48.8305 21.013H43.5433V53.7839H48.8305V21.013Z" fill="white"/>
|
||||
<path d="M57.0987 32.7729H51.8115V53.7835H57.0987V32.7729Z" fill="white"/>
|
||||
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
|
||||
<path d="M58.4255 24.7602L54.4794 20.7832L50.5333 24.7602L54.4794 28.7372L58.4255 24.7602Z" fill="white"/>
|
||||
<path d="M40.5557 21.0118H35.2685V24.185C33.9942 23.5456 32.5588 23.1832 31.0387 23.1832C25.7911 23.1806 21.5217 27.4834 21.5217 32.7721C21.5217 35.082 22.336 37.2054 23.6921 38.8626H21.5217V44.1912H26.8089V41.3618C28.0832 42.0012 29.5186 42.3635 31.0387 42.3635C36.2863 42.3635 40.5557 38.0607 40.5557 32.7721C40.5557 30.2996 39.6225 28.0429 38.0918 26.3404H40.5557V21.0118ZM31.0387 37.0349C28.707 37.0349 26.8089 35.122 26.8089 32.7721C26.8089 30.4221 28.707 28.5092 31.0387 28.5092C33.3704 28.5092 35.2685 30.4221 35.2685 32.7721C35.2685 35.122 33.3704 37.0349 31.0387 37.0349Z" fill="white"/>
|
||||
<path d="M31.0381 44.1912H26.8083V49.5198H31.0381C33.3697 49.5198 35.2678 51.4327 35.2678 53.7826H40.555C40.555 48.494 36.2856 44.1912 31.0381 44.1912Z" fill="white"/>
|
||||
<path d="M69.6041 26.3416C71.1295 26.3416 72.6073 27.1702 73.3871 28.6968L73.596 28.4864L77.2098 24.8443C75.4253 22.4544 72.6126 21.013 69.6041 21.013C64.4438 21.013 60.2352 25.172 60.0951 30.338L60.0872 53.7839H65.3744V30.6018C65.3744 28.2519 67.2725 26.3389 69.6041 26.3389V26.3416Z" fill="white"/>
|
||||
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
|
||||
<path d="M120.022 53.8259H117.218V32.6354H120.022V35.2219C121.02 33.321 122.702 32.2615 125.102 32.2615C129.371 32.2615 131.397 35.9698 131.397 40.5819C131.397 45.1939 129.371 48.9023 125.102 48.9023C122.702 48.9023 121.02 47.8427 120.022 45.9418V53.8259ZM120.022 39.2419V41.9219C120.022 44.6642 121.581 46.6586 124.385 46.6586C126.722 46.6586 128.436 45.1939 128.436 43.0437V38.12C128.436 35.9698 126.722 34.5052 124.385 34.5052C121.581 34.5052 120.022 36.4996 120.022 39.2419Z" fill="white"/>
|
||||
<path d="M103.267 53.8259H100.463V32.6354H103.267V35.2219C104.265 33.321 105.947 32.2615 108.347 32.2615C112.616 32.2615 114.642 35.9698 114.642 40.5819C114.642 45.1939 112.616 48.9023 108.347 48.9023C105.947 48.9023 104.265 47.8427 103.267 45.9418V53.8259ZM103.267 39.2419V41.9219C103.267 44.6642 104.826 46.6586 107.63 46.6586C109.967 46.6586 111.681 45.1939 111.681 43.0437V38.12C111.681 35.9698 109.967 34.5052 107.63 34.5052C104.826 34.5052 103.267 36.4996 103.267 39.2419Z" fill="white"/>
|
||||
<path d="M87.7844 48.9023C86.2262 48.9023 84.8862 48.4037 83.9825 47.4688C83.1723 46.6274 82.6737 45.4121 82.6737 44.1656C82.6737 41.6726 84.263 39.6782 87.4104 39.3977L92.5834 38.9303V37.6526C92.5834 35.1907 91.2434 34.5052 89.2802 34.5052C87.3169 34.5052 86.0081 35.3466 86.0081 37.3721H83.2035C83.2035 34.3805 85.9146 32.2615 89.3113 32.2615C92.7392 32.2615 95.3257 33.9442 95.3257 37.6838V46.1288H97.694V48.5283H92.9573L92.895 45.599H92.7704C91.8978 47.8427 90.1527 48.9023 87.7844 48.9023ZM88.5011 46.6586C91.0253 46.6586 92.5834 44.6642 92.5834 41.7972V41.0805L86.943 41.6102C85.79 41.7037 85.6341 42.2335 85.6341 43.1372V44.5395C85.6341 46.0041 86.7248 46.6586 88.5011 46.6586Z" fill="white"/>
|
||||
<path d="M80.5279 45.1002V48.5281H77.1V45.1002H80.5279Z" fill="white"/>
|
||||
</g>
|
||||
<path d="M102.683 12.3401L100.874 9.09521H101.922L102.976 11.0436C103.015 11.1169 103.05 11.1852 103.079 11.2487C103.108 11.3122 103.138 11.3757 103.167 11.4392C103.191 11.3952 103.211 11.3537 103.225 11.3147C103.24 11.2756 103.257 11.2365 103.277 11.1975C103.301 11.1535 103.328 11.1022 103.357 11.0436L104.405 9.09521H105.423L103.621 12.3401V14.4497H102.683V12.3401Z" fill="white"/>
|
||||
<path d="M97.1749 9.09521V14.4497H96.2373V9.09521H97.1749ZM98.3615 12.1717H96.8892V11.3806H98.3102C98.5691 11.3806 98.7668 11.3171 98.9036 11.1901C99.0403 11.0583 99.1087 10.8727 99.1087 10.6334C99.1087 10.4039 99.0378 10.2281 98.8962 10.106C98.7546 9.98397 98.5495 9.92292 98.2809 9.92292H96.8599V9.09521H98.3615C98.8938 9.09521 99.3113 9.22462 99.6141 9.48343C99.9168 9.74224 100.068 10.0963 100.068 10.5455C100.068 10.8678 99.9901 11.1389 99.8338 11.3586C99.6776 11.5735 99.4456 11.7297 99.1379 11.8274V11.7248C99.47 11.803 99.7215 11.9495 99.8924 12.1643C100.063 12.3792 100.149 12.6575 100.149 12.9994C100.149 13.3021 100.08 13.5634 99.9437 13.7831C99.807 13.998 99.6067 14.164 99.343 14.2812C99.0842 14.3935 98.7717 14.4497 98.4055 14.4497H96.8599V13.622H98.3615C98.6301 13.622 98.8352 13.5585 98.9768 13.4315C99.1184 13.3046 99.1892 13.1215 99.1892 12.8822C99.1892 12.6575 99.116 12.4842 98.9695 12.3621C98.8279 12.2351 98.6252 12.1717 98.3615 12.1717Z" fill="white"/>
|
||||
<path d="M89.5954 14.4497H87.6689V9.09521H89.5441C90.0715 9.09521 90.5354 9.20997 90.9358 9.43948C91.3363 9.66411 91.6488 9.97908 91.8734 10.3844C92.1029 10.7848 92.2177 11.2512 92.2177 11.7834C92.2177 12.3059 92.1054 12.7699 91.8807 13.1752C91.661 13.5756 91.3534 13.8881 90.9578 14.1128C90.5671 14.3374 90.113 14.4497 89.5954 14.4497ZM88.6065 9.52738V14.0249L88.1597 13.5854H89.5075C89.864 13.5854 90.1716 13.5121 90.4304 13.3656C90.6892 13.2191 90.887 13.0116 91.0237 12.743C91.1605 12.4744 91.2288 12.1546 91.2288 11.7834C91.2288 11.4025 91.158 11.0778 91.0164 10.8092C90.8748 10.5358 90.6721 10.3258 90.4084 10.1793C90.1447 10.0328 89.8273 9.95955 89.4562 9.95955H88.1597L88.6065 9.52738Z" fill="white"/>
|
||||
<path d="M86.0735 14.4497H82.748V9.09521H86.0735V9.95955H83.356L83.6856 9.65923V11.3366H85.8245V12.1643H83.6856V13.8857L83.356 13.5854H86.0735V14.4497Z" fill="white"/>
|
||||
<path d="M78.1926 14.4497H77.255V9.09521H79.2986C79.9042 9.09521 80.3754 9.24171 80.7123 9.53471C81.0542 9.8277 81.2251 10.2379 81.2251 10.7653C81.2251 11.1218 81.1421 11.427 80.976 11.6809C80.8149 11.9299 80.5756 12.1204 80.2582 12.2522L81.2764 14.4497H80.2509L79.3426 12.45H78.1926V14.4497ZM78.1926 9.93025V11.6223H79.2986C79.5965 11.6223 79.8285 11.5466 79.9945 11.3952C80.1605 11.2438 80.2436 11.0339 80.2436 10.7653C80.2436 10.4967 80.1605 10.2916 79.9945 10.15C79.8285 10.0035 79.5965 9.93025 79.2986 9.93025H78.1926Z" fill="white"/>
|
||||
<path d="M75.789 11.7688C75.789 12.3108 75.6792 12.7918 75.4594 13.2118C75.2397 13.6269 74.9345 13.9516 74.5438 14.186C74.1531 14.4204 73.7014 14.5376 73.1887 14.5376C72.6808 14.5376 72.2316 14.4204 71.8409 14.186C71.4503 13.9516 71.1451 13.6269 70.9253 13.2118C70.7105 12.7967 70.603 12.3182 70.603 11.7761C70.603 11.2292 70.7129 10.7482 70.9326 10.3331C71.1524 9.91317 71.4552 9.58599 71.8409 9.35159C72.2316 9.1172 72.6833 9 73.196 9C73.7088 9 74.158 9.1172 74.5438 9.35159C74.9345 9.58599 75.2397 9.91073 75.4594 10.3258C75.6792 10.7409 75.789 11.2219 75.789 11.7688ZM74.8075 11.7688C74.8075 11.3879 74.7416 11.0583 74.6097 10.7799C74.4779 10.5016 74.2923 10.2867 74.053 10.1354C73.8138 9.97909 73.5281 9.90096 73.196 9.90096C72.8689 9.90096 72.5832 9.97909 72.339 10.1354C72.0997 10.2867 71.9142 10.5016 71.7823 10.7799C71.6505 11.0583 71.5846 11.3879 71.5846 11.7688C71.5846 12.1497 71.6505 12.4818 71.7823 12.765C71.9142 13.0433 72.0997 13.2582 72.339 13.4096C72.5832 13.561 72.8689 13.6366 73.196 13.6366C73.5281 13.6366 73.8138 13.561 74.053 13.4096C74.2923 13.2533 74.4779 13.036 74.6097 12.7577C74.7416 12.4744 74.8075 12.1448 74.8075 11.7688Z" fill="white"/>
|
||||
<path d="M65.6821 10.5895C65.6821 10.277 65.7627 10.0011 65.9239 9.76179C66.085 9.52251 66.3072 9.33694 66.5904 9.2051C66.8785 9.06837 67.2106 9 67.5866 9C67.948 9 68.2605 9.06348 68.5242 9.19045C68.7928 9.31741 69.0003 9.49809 69.1468 9.73249C69.2982 9.96688 69.3788 10.2452 69.3885 10.5675H68.4509C68.4412 10.338 68.3582 10.1598 68.2019 10.0328C68.0456 9.90096 67.8356 9.83504 67.572 9.83504C67.2838 9.83504 67.0519 9.90096 66.8761 10.0328C66.7052 10.1598 66.6197 10.3356 66.6197 10.5602C66.6197 10.7506 66.671 10.902 66.7735 11.0143C66.881 11.1218 67.047 11.2023 67.2716 11.2561L68.114 11.4465C68.573 11.5442 68.9148 11.7126 69.1395 11.9519C69.3641 12.1863 69.4764 12.5037 69.4764 12.9042C69.4764 13.2313 69.3958 13.5194 69.2347 13.7685C69.0736 14.0175 68.844 14.2104 68.5462 14.3472C68.2532 14.479 67.9089 14.5449 67.5134 14.5449C67.1373 14.5449 66.8077 14.4814 66.5245 14.3545C66.2413 14.2226 66.0191 14.0395 65.8579 13.8051C65.7017 13.5707 65.6187 13.2948 65.6089 12.9774H66.5465C66.5514 13.202 66.6393 13.3803 66.8102 13.5121C66.986 13.6391 67.2228 13.7026 67.5207 13.7026C67.8332 13.7026 68.0798 13.6391 68.2605 13.5121C68.4461 13.3803 68.5388 13.2069 68.5388 12.9921C68.5388 12.8065 68.49 12.66 68.3923 12.5526C68.2947 12.4402 68.136 12.3621 67.9162 12.3182L67.0665 12.1277C66.6124 12.0301 66.2681 11.8543 66.0337 11.6003C65.7993 11.3415 65.6821 11.0046 65.6821 10.5895Z" fill="white"/>
|
||||
<path d="M60.8331 14.4497H59.9102V9.09521H60.8404L63.6239 13.307H63.3528V9.09521H64.2758V14.4497H63.3528L60.5621 10.2452H60.8331V14.4497Z" fill="white"/>
|
||||
<path d="M58.4443 11.7688C58.4443 12.3108 58.3344 12.7918 58.1147 13.2118C57.8949 13.6269 57.5897 13.9516 57.1991 14.186C56.8084 14.4204 56.3567 14.5376 55.844 14.5376C55.3361 14.5376 54.8869 14.4204 54.4962 14.186C54.1055 13.9516 53.8003 13.6269 53.5806 13.2118C53.3657 12.7967 53.2583 12.3182 53.2583 11.7761C53.2583 11.2292 53.3682 10.7482 53.5879 10.3331C53.8077 9.91317 54.1104 9.58599 54.4962 9.35159C54.8869 9.1172 55.3386 9 55.8513 9C56.364 9 56.8133 9.1172 57.1991 9.35159C57.5897 9.58599 57.8949 9.91073 58.1147 10.3258C58.3344 10.7409 58.4443 11.2219 58.4443 11.7688ZM57.4628 11.7688C57.4628 11.3879 57.3969 11.0583 57.265 10.7799C57.1332 10.5016 56.9476 10.2867 56.7083 10.1354C56.469 9.97909 56.1834 9.90096 55.8513 9.90096C55.5241 9.90096 55.2385 9.97909 54.9943 10.1354C54.755 10.2867 54.5695 10.5016 54.4376 10.7799C54.3058 11.0583 54.2398 11.3879 54.2398 11.7688C54.2398 12.1497 54.3058 12.4818 54.4376 12.765C54.5695 13.0433 54.755 13.2582 54.9943 13.4096C55.2385 13.561 55.5241 13.6366 55.8513 13.6366C56.1834 13.6366 56.469 13.561 56.7083 13.4096C56.9476 13.2533 57.1332 13.036 57.265 12.7577C57.3969 12.4744 57.4628 12.1448 57.4628 11.7688Z" fill="white"/>
|
||||
<path d="M49.2393 9.09521V14.4497H48.3018V9.09521H49.2393ZM50.4186 12.6038H49.0123V11.7688H50.2209C50.5432 11.7688 50.7873 11.6882 50.9534 11.5271C51.1243 11.361 51.2097 11.1315 51.2097 10.8385C51.2097 10.5455 51.1243 10.3209 50.9534 10.1646C50.7873 10.0084 50.5481 9.93025 50.2355 9.93025H48.9244V9.09521H50.4186C50.78 9.09521 51.0925 9.16846 51.3562 9.31496C51.6199 9.46146 51.825 9.66655 51.9715 9.93025C52.118 10.1891 52.1913 10.4943 52.1913 10.8459C52.1913 11.1877 52.118 11.4929 51.9715 11.7615C51.825 12.0252 51.6199 12.2327 51.3562 12.3841C51.0925 12.5306 50.78 12.6038 50.4186 12.6038Z" fill="white"/>
|
||||
<path d="M43.0732 10.5895C43.0732 10.277 43.1538 10.0011 43.315 9.76179C43.4761 9.52251 43.6983 9.33694 43.9815 9.2051C44.2696 9.06837 44.6017 9 44.9777 9C45.3391 9 45.6516 9.06348 45.9153 9.19045C46.1839 9.31741 46.3914 9.49809 46.5379 9.73249C46.6893 9.96688 46.7699 10.2452 46.7796 10.5675H45.842C45.8323 10.338 45.7493 10.1598 45.593 10.0328C45.4367 9.90096 45.2268 9.83504 44.9631 9.83504C44.675 9.83504 44.443 9.90096 44.2672 10.0328C44.0963 10.1598 44.0108 10.3356 44.0108 10.5602C44.0108 10.7506 44.0621 10.902 44.1647 11.0143C44.2721 11.1218 44.4381 11.2023 44.6627 11.2561L45.5051 11.4465C45.9641 11.5442 46.306 11.7126 46.5306 11.9519C46.7552 12.1863 46.8675 12.5037 46.8675 12.9042C46.8675 13.2313 46.787 13.5194 46.6258 13.7685C46.4647 14.0175 46.2352 14.2104 45.9373 14.3472C45.6443 14.479 45.3 14.5449 44.9045 14.5449C44.5285 14.5449 44.1988 14.4814 43.9156 14.3545C43.6324 14.2226 43.4102 14.0395 43.249 13.8051C43.0928 13.5707 43.0098 13.2948 43 12.9774H43.9376C43.9425 13.202 44.0304 13.3803 44.2013 13.5121C44.3771 13.6391 44.6139 13.7026 44.9118 13.7026C45.2243 13.7026 45.4709 13.6391 45.6516 13.5121C45.8372 13.3803 45.9299 13.2069 45.9299 12.9921C45.9299 12.8065 45.8811 12.66 45.7835 12.5526C45.6858 12.4402 45.5271 12.3621 45.3073 12.3182L44.4576 12.1277C44.0035 12.0301 43.6592 11.8543 43.4248 11.6003C43.1904 11.3415 43.0732 11.0046 43.0732 10.5895Z" fill="white"/>
|
||||
<defs>
|
||||
<filter id="filter0_d_10631_12135" x="20.0869" y="19.3484" width="112.745" height="35.9123" filterUnits="userSpaceOnUse" color-interpolation-filters="sRGB">
|
||||
<feFlood flood-opacity="0" result="BackgroundImageFix"/>
|
||||
<feColorMatrix in="SourceAlpha" type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 127 0" result="hardAlpha"/>
|
||||
<feOffset/>
|
||||
<feGaussianBlur stdDeviation="0.717391"/>
|
||||
<feComposite in2="hardAlpha" operator="out"/>
|
||||
<feColorMatrix type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0.25 0"/>
|
||||
<feBlend mode="normal" in2="BackgroundImageFix" result="effect1_dropShadow_10631_12135"/>
|
||||
<feBlend mode="normal" in="SourceGraphic" in2="effect1_dropShadow_10631_12135" result="shape"/>
|
||||
</filter>
|
||||
</defs>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 13 KiB |
8
build_and_push_docker.yaml
Normal file
8
build_and_push_docker.yaml
Normal file
@@ -0,0 +1,8 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
|
||||
# wait 2 seconds
|
||||
sleep 2
|
||||
docker build --build-arg CACHEBUST=$(date +%s) -t aitoolkit:latest -f docker/Dockerfile .
|
||||
docker tag aitoolkit:latest ostris/aitoolkit:latest
|
||||
docker push ostris/aitoolkit:latest
|
||||
96
config/examples/modal/modal_train_lora_flux_24gb.yaml
Normal file
96
config/examples/modal/modal_train_lora_flux_24gb.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
|
||||
- folder_path: "/root/ai-toolkit/your-dataset"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
# if you get an error, or get stuck while downloading,
|
||||
# check https://github.com/ostris/ai-toolkit/issues/84, download the model locally and
|
||||
# place it like "/root/ai-toolkit/FLUX.1-dev"
|
||||
name_or_path: "black-forest-labs/FLUX.1-dev"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 20
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
@@ -0,0 +1,98 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
|
||||
- folder_path: "/root/ai-toolkit/your-dataset"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
# if you get an error, or get stuck while downloading,
|
||||
# check https://github.com/ostris/ai-toolkit/issues/84, download the models locally and
|
||||
# place them like "/root/ai-toolkit/FLUX.1-schnell" and "/root/ai-toolkit/FLUX.1-schnell-training-adapter"
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
96
config/examples/train_lora_flux_24gb.yaml
Normal file
96
config/examples/train_lora_flux_24gb.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "black-forest-labs/FLUX.1-dev"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 20
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
98
config/examples/train_lora_flux_schnell_24gb.yaml
Normal file
98
config/examples/train_lora_flux_schnell_24gb.yaml
Normal file
@@ -0,0 +1,98 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flux_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flux
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new bell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "black-forest-labs/FLUX.1-schnell"
|
||||
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
|
||||
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flux
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 1 # schnell does not do guidance
|
||||
sample_steps: 4 # 1 - 4 works well
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
21
docker/Dockerfile
Normal file
21
docker/Dockerfile
Normal file
@@ -0,0 +1,21 @@
|
||||
FROM runpod/base:0.6.2-cuda12.1.0
|
||||
LABEL authors="jaret"
|
||||
|
||||
# Install dependencies
|
||||
RUN apt-get update
|
||||
|
||||
WORKDIR /app
|
||||
ARG CACHEBUST=1
|
||||
RUN git clone https://github.com/ostris/ai-toolkit.git && \
|
||||
cd ai-toolkit && \
|
||||
git submodule update --init --recursive
|
||||
|
||||
WORKDIR /app/ai-toolkit
|
||||
|
||||
RUN ln -s /usr/bin/python3 /usr/bin/python
|
||||
RUN python -m pip install -r requirements.txt
|
||||
|
||||
RUN apt-get install -y tmux nvtop htop
|
||||
|
||||
WORKDIR /
|
||||
CMD ["/start.sh"]
|
||||
256
extensions_built_in/advanced_generator/Img2ImgGenerator.py
Normal file
256
extensions_built_in/advanced_generator/Img2ImgGenerator.py
Normal file
@@ -0,0 +1,256 @@
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import List
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLImg2ImgPipeline, PixArtSigmaPipeline
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
from toolkit.sampler import get_sampler
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseExtensionProcess
|
||||
from toolkit.data_loader import get_dataloader_from_datasets
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
from controlnet_aux.midas import MidasDetector
|
||||
from diffusers.utils import load_image
|
||||
from torchvision.transforms import ToTensor
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class GenerateConfig:
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.prompts: List[str]
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.walk_seed = kwargs.get('walk_seed', False)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
self.sample_steps = kwargs.get('sample_steps', 20)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.denoise_strength = kwargs.get('denoise_strength', 0.5)
|
||||
self.trigger_word = kwargs.get('trigger_word', None)
|
||||
|
||||
|
||||
class Img2ImgGenerator(BaseExtensionProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict):
|
||||
super().__init__(process_id, job, config)
|
||||
self.output_folder = self.get_conf('output_folder', required=True)
|
||||
self.copy_inputs_to = self.get_conf('copy_inputs_to', None)
|
||||
self.device = self.get_conf('device', 'cuda')
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
self.is_latents_cached = True
|
||||
raw_datasets = self.get_conf('datasets', None)
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
raw_datasets = preprocess_dataset_raw_config(raw_datasets)
|
||||
self.datasets = None
|
||||
self.datasets_reg = None
|
||||
self.dtype = self.get_conf('dtype', 'float16')
|
||||
self.torch_dtype = get_torch_dtype(self.dtype)
|
||||
self.params = []
|
||||
if raw_datasets is not None and len(raw_datasets) > 0:
|
||||
for raw_dataset in raw_datasets:
|
||||
dataset = DatasetConfig(**raw_dataset)
|
||||
is_caching = dataset.cache_latents or dataset.cache_latents_to_disk
|
||||
if not is_caching:
|
||||
self.is_latents_cached = False
|
||||
if dataset.is_reg:
|
||||
if self.datasets_reg is None:
|
||||
self.datasets_reg = []
|
||||
self.datasets_reg.append(dataset)
|
||||
else:
|
||||
if self.datasets is None:
|
||||
self.datasets = []
|
||||
self.datasets.append(dataset)
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
device=self.device,
|
||||
model_config=self.model_config,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
print(f"Using device {self.device}")
|
||||
self.data_loader: DataLoader = None
|
||||
self.adapter: T2IAdapter = None
|
||||
|
||||
def to_pil(self, img):
|
||||
# image comes in -1 to 1. convert to a PIL RGB image
|
||||
img = (img + 1) / 2
|
||||
img = img.clamp(0, 1)
|
||||
img = img[0].permute(1, 2, 0).cpu().numpy()
|
||||
img = (img * 255).astype(np.uint8)
|
||||
image = Image.fromarray(img)
|
||||
return image
|
||||
|
||||
def run(self):
|
||||
with torch.no_grad():
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
device = torch.device(self.device)
|
||||
|
||||
if self.model_config.is_xl:
|
||||
pipe = StableDiffusionXLImg2ImgPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
elif self.model_config.is_pixart:
|
||||
pipe = self.sd.pipeline.to(device, dtype=self.torch_dtype)
|
||||
else:
|
||||
raise NotImplementedError("Only XL models are supported")
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
|
||||
# pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
|
||||
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
|
||||
|
||||
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
|
||||
|
||||
num_batches = len(self.data_loader)
|
||||
pbar = tqdm(total=num_batches, desc="Generating images")
|
||||
seed = self.generate_config.seed
|
||||
# load images from datasets, use tqdm
|
||||
for i, batch in enumerate(self.data_loader):
|
||||
batch: DataLoaderBatchDTO = batch
|
||||
|
||||
gen_seed = seed if seed > 0 else random.randint(0, 2 ** 32 - 1)
|
||||
generator = torch.manual_seed(gen_seed)
|
||||
|
||||
file_item: FileItemDTO = batch.file_items[0]
|
||||
img_path = file_item.path
|
||||
img_filename = os.path.basename(img_path)
|
||||
img_filename_no_ext = os.path.splitext(img_filename)[0]
|
||||
img_filename = img_filename_no_ext + '.' + self.generate_config.ext
|
||||
output_path = os.path.join(self.output_folder, img_filename)
|
||||
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
|
||||
|
||||
if self.copy_inputs_to is not None:
|
||||
output_inputs_path = os.path.join(self.copy_inputs_to, img_filename)
|
||||
output_inputs_caption_path = os.path.join(self.copy_inputs_to, img_filename_no_ext + '.txt')
|
||||
else:
|
||||
output_inputs_path = None
|
||||
output_inputs_caption_path = None
|
||||
|
||||
caption = batch.get_caption_list()[0]
|
||||
if self.generate_config.trigger_word is not None:
|
||||
caption = caption.replace('[trigger]', self.generate_config.trigger_word)
|
||||
|
||||
img: torch.Tensor = batch.tensor.clone()
|
||||
image = self.to_pil(img)
|
||||
|
||||
# image.save(output_depth_path)
|
||||
if self.model_config.is_pixart:
|
||||
pipe: PixArtSigmaPipeline = pipe
|
||||
|
||||
# Encode the full image once
|
||||
encoded_image = pipe.vae.encode(
|
||||
pipe.image_processor.preprocess(image).to(device=pipe.device, dtype=pipe.dtype))
|
||||
if hasattr(encoded_image, "latent_dist"):
|
||||
latents = encoded_image.latent_dist.sample(generator)
|
||||
elif hasattr(encoded_image, "latents"):
|
||||
latents = encoded_image.latents
|
||||
else:
|
||||
raise AttributeError("Could not access latents of provided encoder_output")
|
||||
latents = pipe.vae.config.scaling_factor * latents
|
||||
|
||||
# latents = self.sd.encode_images(img)
|
||||
|
||||
# self.sd.noise_scheduler.set_timesteps(self.generate_config.sample_steps)
|
||||
# start_step = math.floor(self.generate_config.sample_steps * self.generate_config.denoise_strength)
|
||||
# timestep = self.sd.noise_scheduler.timesteps[start_step].unsqueeze(0)
|
||||
# timestep = timestep.to(device, dtype=torch.int32)
|
||||
# latent = latent.to(device, dtype=self.torch_dtype)
|
||||
# noise = torch.randn_like(latent, device=device, dtype=self.torch_dtype)
|
||||
# latent = self.sd.add_noise(latent, noise, timestep)
|
||||
# timesteps_to_use = self.sd.noise_scheduler.timesteps[start_step + 1:]
|
||||
batch_size = 1
|
||||
num_images_per_prompt = 1
|
||||
|
||||
shape = (batch_size, pipe.transformer.config.in_channels, image.height // pipe.vae_scale_factor,
|
||||
image.width // pipe.vae_scale_factor)
|
||||
noise = randn_tensor(shape, generator=generator, device=pipe.device, dtype=pipe.dtype)
|
||||
|
||||
# noise = torch.randn_like(latents, device=device, dtype=self.torch_dtype)
|
||||
num_inference_steps = self.generate_config.sample_steps
|
||||
strength = self.generate_config.denoise_strength
|
||||
# Get timesteps
|
||||
init_timestep = min(int(num_inference_steps * strength), num_inference_steps)
|
||||
t_start = max(num_inference_steps - init_timestep, 0)
|
||||
pipe.scheduler.set_timesteps(num_inference_steps, device="cpu")
|
||||
timesteps = pipe.scheduler.timesteps[t_start:]
|
||||
timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
|
||||
latents = pipe.scheduler.add_noise(latents, noise, timestep)
|
||||
|
||||
gen_images = pipe.__call__(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
latents=latents,
|
||||
timesteps=timesteps,
|
||||
width=image.width,
|
||||
height=image.height,
|
||||
num_inference_steps=num_inference_steps,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
# strength=self.generate_config.denoise_strength,
|
||||
use_resolution_binning=False,
|
||||
output_type="np"
|
||||
).images[0]
|
||||
gen_images = (gen_images * 255).clip(0, 255).astype(np.uint8)
|
||||
gen_images = Image.fromarray(gen_images)
|
||||
else:
|
||||
pipe: StableDiffusionXLImg2ImgPipeline = pipe
|
||||
|
||||
gen_images = pipe.__call__(
|
||||
prompt=caption,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
image=image,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
strength=self.generate_config.denoise_strength,
|
||||
).images[0]
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
gen_images.save(output_path)
|
||||
|
||||
# save caption
|
||||
with open(output_caption_path, 'w') as f:
|
||||
f.write(caption)
|
||||
|
||||
if output_inputs_path is not None:
|
||||
os.makedirs(os.path.dirname(output_inputs_path), exist_ok=True)
|
||||
image.save(output_inputs_path)
|
||||
with open(output_inputs_caption_path, 'w') as f:
|
||||
f.write(caption)
|
||||
|
||||
pbar.update(1)
|
||||
batch.cleanup()
|
||||
|
||||
pbar.close()
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
@@ -7,7 +7,7 @@ import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers import T2IAdapter
|
||||
from torch.utils.data import DataLoader
|
||||
from diffusers import StableDiffusionXLAdapterPipeline
|
||||
from diffusers import StableDiffusionXLAdapterPipeline, StableDiffusionAdapterPipeline
|
||||
from tqdm import tqdm
|
||||
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
|
||||
@@ -100,25 +100,43 @@ class ReferenceGenerator(BaseExtensionProcess):
|
||||
|
||||
if self.generate_config.t2i_adapter_path is not None:
|
||||
self.adapter = T2IAdapter.from_pretrained(
|
||||
"TencentARC/t2i-adapter-depth-midas-sdxl-1.0", torch_dtype=self.torch_dtype, varient="fp16"
|
||||
self.generate_config.t2i_adapter_path,
|
||||
torch_dtype=self.torch_dtype,
|
||||
varient="fp16"
|
||||
).to(device)
|
||||
|
||||
midas_depth = MidasDetector.from_pretrained(
|
||||
"valhalla/t2iadapter-aux-models", filename="dpt_large_384.pt", model_type="dpt_large"
|
||||
).to(device)
|
||||
|
||||
pipe = StableDiffusionXLAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
adapter=self.adapter,
|
||||
).to(device)
|
||||
if self.model_config.is_xl:
|
||||
pipe = StableDiffusionXLAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder[0],
|
||||
text_encoder_2=self.sd.text_encoder[1],
|
||||
tokenizer=self.sd.tokenizer[0],
|
||||
tokenizer_2=self.sd.tokenizer[1],
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
adapter=self.adapter,
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
else:
|
||||
pipe = StableDiffusionAdapterPipeline(
|
||||
vae=self.sd.vae,
|
||||
unet=self.sd.unet,
|
||||
text_encoder=self.sd.text_encoder,
|
||||
tokenizer=self.sd.tokenizer,
|
||||
scheduler=get_sampler(self.generate_config.sampler),
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
adapter=self.adapter,
|
||||
).to(device, dtype=self.torch_dtype)
|
||||
pipe.set_progress_bar_config(disable=True)
|
||||
|
||||
pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
|
||||
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
|
||||
|
||||
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
|
||||
|
||||
num_batches = len(self.data_loader)
|
||||
@@ -176,6 +194,7 @@ class ReferenceGenerator(BaseExtensionProcess):
|
||||
adapter_conditioning_scale=self.generate_config.adapter_conditioning_scale,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
).images[0]
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
gen_images.save(output_path)
|
||||
|
||||
# save caption
|
||||
|
||||
@@ -36,7 +36,24 @@ class PureLoraGenerator(Extension):
|
||||
return PureLoraGenerator
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class Img2ImgGeneratorExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "batch_img2img"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Img2ImgGeneratorExtension"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .Img2ImgGenerator import Img2ImgGenerator
|
||||
return Img2ImgGenerator
|
||||
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
AdvancedReferenceGeneratorExtension, PureLoraGenerator
|
||||
AdvancedReferenceGeneratorExtension, PureLoraGenerator, Img2ImgGeneratorExtension
|
||||
]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
249
extensions_built_in/sd_trainer/TrainerV2.py
Normal file
249
extensions_built_in/sd_trainer/TrainerV2.py
Normal file
@@ -0,0 +1,249 @@
|
||||
import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import Union, List
|
||||
|
||||
import numpy as np
|
||||
from diffusers import T2IAdapter, ControlNetModel
|
||||
import torch.distributed as dist
|
||||
from torch import nn
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
from toolkit.clip_vision_adapter import ClipVisionAdapter
|
||||
from toolkit.data_loader import get_dataloader_datasets
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.stable_diffusion_model import BlankNetwork
|
||||
from toolkit.train_tools import get_torch_dtype, add_all_snr_to_noise_scheduler
|
||||
import gc
|
||||
import torch
|
||||
from jobs.process import BaseSDTrainProcess
|
||||
from torchvision import transforms
|
||||
from diffusers import EMAModel
|
||||
import math
|
||||
from toolkit.train_tools import precondition_model_outputs_flow_match
|
||||
from toolkit.models.unified_training_model import UnifiedTrainingModel
|
||||
|
||||
|
||||
def flush():
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
adapter_transforms = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
|
||||
|
||||
class TrainerV2(BaseSDTrainProcess):
|
||||
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super().__init__(process_id, job, config, **kwargs)
|
||||
self.assistant_adapter: Union['T2IAdapter', 'ControlNetModel', None]
|
||||
self.do_prior_prediction = False
|
||||
self.do_long_prompts = False
|
||||
self.do_guided_loss = False
|
||||
|
||||
self._clip_image_embeds_unconditional: Union[List[str], None] = None
|
||||
self.negative_prompt_pool: Union[List[str], None] = None
|
||||
self.batch_negative_prompt: Union[List[str], None] = None
|
||||
|
||||
self.scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
self.is_bfloat = self.train_config.dtype == "bfloat16" or self.train_config.dtype == "bf16"
|
||||
|
||||
self.do_grad_scale = True
|
||||
if self.is_fine_tuning:
|
||||
self.do_grad_scale = False
|
||||
if self.adapter_config is not None:
|
||||
if self.adapter_config.train:
|
||||
self.do_grad_scale = False
|
||||
|
||||
if self.train_config.dtype in ["fp16", "float16"]:
|
||||
# patch the scaler to allow fp16 training
|
||||
org_unscale_grads = self.scaler._unscale_grads_
|
||||
def _unscale_grads_replacer(optimizer, inv_scale, found_inf, allow_fp16):
|
||||
return org_unscale_grads(optimizer, inv_scale, found_inf, True)
|
||||
self.scaler._unscale_grads_ = _unscale_grads_replacer
|
||||
|
||||
self.unified_training_model: UnifiedTrainingModel = None
|
||||
self.device_ids = list(range(torch.cuda.device_count()))
|
||||
|
||||
def before_model_load(self):
|
||||
pass
|
||||
|
||||
def before_dataset_load(self):
|
||||
self.assistant_adapter = None
|
||||
# get adapter assistant if one is set
|
||||
if self.train_config.adapter_assist_name_or_path is not None:
|
||||
adapter_path = self.train_config.adapter_assist_name_or_path
|
||||
|
||||
if self.train_config.adapter_assist_type == "t2i":
|
||||
# dont name this adapter since we are not training it
|
||||
self.assistant_adapter = T2IAdapter.from_pretrained(
|
||||
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype)
|
||||
).to(self.device_torch)
|
||||
elif self.train_config.adapter_assist_type == "control_net":
|
||||
self.assistant_adapter = ControlNetModel.from_pretrained(
|
||||
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype)
|
||||
).to(self.device_torch, dtype=get_torch_dtype(self.train_config.dtype))
|
||||
else:
|
||||
raise ValueError(f"Unknown adapter assist type {self.train_config.adapter_assist_type}")
|
||||
|
||||
self.assistant_adapter.eval()
|
||||
self.assistant_adapter.requires_grad_(False)
|
||||
flush()
|
||||
if self.train_config.train_turbo and self.train_config.show_turbo_outputs:
|
||||
raise ValueError("Turbo outputs are not supported on MultiGPUSDTrainer")
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
# if self.train_config.do_prior_divergence:
|
||||
# self.do_prior_prediction = True
|
||||
# move vae to device if we did not cache latents
|
||||
if not self.is_latents_cached:
|
||||
self.sd.vae.eval()
|
||||
self.sd.vae.to(self.device_torch)
|
||||
else:
|
||||
# offload it. Already cached
|
||||
self.sd.vae.to('cpu')
|
||||
flush()
|
||||
add_all_snr_to_noise_scheduler(self.sd.noise_scheduler, self.device_torch)
|
||||
if self.adapter is not None:
|
||||
self.adapter.to(self.device_torch)
|
||||
|
||||
# check if we have regs and using adapter and caching clip embeddings
|
||||
has_reg = self.datasets_reg is not None and len(self.datasets_reg) > 0
|
||||
is_caching_clip_embeddings = self.datasets is not None and any([self.datasets[i].cache_clip_vision_to_disk for i in range(len(self.datasets))])
|
||||
|
||||
if has_reg and is_caching_clip_embeddings:
|
||||
# we need a list of unconditional clip image embeds from other datasets to handle regs
|
||||
unconditional_clip_image_embeds = []
|
||||
datasets = get_dataloader_datasets(self.data_loader)
|
||||
for i in range(len(datasets)):
|
||||
unconditional_clip_image_embeds += datasets[i].clip_vision_unconditional_cache
|
||||
|
||||
if len(unconditional_clip_image_embeds) == 0:
|
||||
raise ValueError("No unconditional clip image embeds found. This should not happen")
|
||||
|
||||
self._clip_image_embeds_unconditional = unconditional_clip_image_embeds
|
||||
|
||||
if self.train_config.negative_prompt is not None:
|
||||
raise ValueError("Negative prompt is not supported on MultiGPUSDTrainer")
|
||||
|
||||
# setup the unified training model
|
||||
self.unified_training_model = UnifiedTrainingModel(
|
||||
sd=self.sd,
|
||||
network=self.network,
|
||||
adapter=self.adapter,
|
||||
assistant_adapter=self.assistant_adapter,
|
||||
train_config=self.train_config,
|
||||
adapter_config=self.adapter_config,
|
||||
embedding=self.embedding,
|
||||
timer=self.timer,
|
||||
trigger_word=self.trigger_word,
|
||||
gpu_ids=self.device_ids,
|
||||
)
|
||||
self.unified_training_model = nn.DataParallel(
|
||||
self.unified_training_model,
|
||||
device_ids=self.device_ids
|
||||
)
|
||||
|
||||
self.unified_training_model = self.unified_training_model.to(self.device_torch)
|
||||
|
||||
# call parent hook
|
||||
super().hook_before_train_loop()
|
||||
|
||||
# you can expand these in a child class to make customization easier
|
||||
|
||||
def preprocess_batch(self, batch: 'DataLoaderBatchDTO'):
|
||||
return self.unified_training_model.preprocess_batch(batch)
|
||||
|
||||
|
||||
def before_unet_predict(self):
|
||||
pass
|
||||
|
||||
def after_unet_predict(self):
|
||||
pass
|
||||
|
||||
def end_of_training_loop(self):
|
||||
pass
|
||||
|
||||
def hook_train_loop(self, batch: 'DataLoaderBatchDTO'):
|
||||
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
loss = self.unified_training_model(batch)
|
||||
|
||||
if torch.isnan(loss):
|
||||
print("loss is nan")
|
||||
loss = torch.zeros_like(loss).requires_grad_(True)
|
||||
|
||||
if self.network is not None:
|
||||
network = self.network
|
||||
else:
|
||||
network = BlankNetwork()
|
||||
|
||||
with (network):
|
||||
with self.timer('backward'):
|
||||
# todo we have multiplier seperated. works for now as res are not in same batch, but need to change
|
||||
# IMPORTANT if gradient checkpointing do not leave with network when doing backward
|
||||
# it will destroy the gradients. This is because the network is a context manager
|
||||
# and will change the multipliers back to 0.0 when exiting. They will be
|
||||
# 0.0 for the backward pass and the gradients will be 0.0
|
||||
# I spent weeks on fighting this. DON'T DO IT
|
||||
# with fsdp_overlap_step_with_backward():
|
||||
# if self.is_bfloat:
|
||||
# loss.backward()
|
||||
# else:
|
||||
if not self.do_grad_scale:
|
||||
loss.backward()
|
||||
else:
|
||||
self.scaler.scale(loss).backward()
|
||||
|
||||
if not self.is_grad_accumulation_step:
|
||||
# fix this for multi params
|
||||
if self.train_config.optimizer != 'adafactor':
|
||||
if self.do_grad_scale:
|
||||
self.scaler.unscale_(self.optimizer)
|
||||
if isinstance(self.params[0], dict):
|
||||
for i in range(len(self.params)):
|
||||
torch.nn.utils.clip_grad_norm_(self.params[i]['params'], self.train_config.max_grad_norm)
|
||||
else:
|
||||
torch.nn.utils.clip_grad_norm_(self.params, self.train_config.max_grad_norm)
|
||||
# only step if we are not accumulating
|
||||
with self.timer('optimizer_step'):
|
||||
# self.optimizer.step()
|
||||
if not self.do_grad_scale:
|
||||
self.optimizer.step()
|
||||
else:
|
||||
self.scaler.step(self.optimizer)
|
||||
self.scaler.update()
|
||||
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
if self.ema is not None:
|
||||
with self.timer('ema_update'):
|
||||
self.ema.update()
|
||||
else:
|
||||
# gradient accumulation. Just a place for breakpoint
|
||||
pass
|
||||
|
||||
# TODO Should we only step scheduler on grad step? If so, need to recalculate last step
|
||||
with self.timer('scheduler_step'):
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if self.embedding is not None:
|
||||
with self.timer('restore_embeddings'):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
self.embedding.restore_embeddings()
|
||||
if self.adapter is not None and isinstance(self.adapter, ClipVisionAdapter):
|
||||
with self.timer('restore_adapter'):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
self.adapter.restore_embeddings()
|
||||
|
||||
loss_dict = OrderedDict(
|
||||
{'loss': loss.item()}
|
||||
)
|
||||
|
||||
self.end_of_training_loop()
|
||||
|
||||
return loss_dict
|
||||
@@ -19,6 +19,23 @@ class SDTrainerExtension(Extension):
|
||||
return SDTrainer
|
||||
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class MultiGPUSDTrainerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "trainer_v2"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "Trainer V2"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .TrainerV2 import TrainerV2
|
||||
return TrainerV2
|
||||
|
||||
|
||||
# for backwards compatability
|
||||
class TextualInversionTrainer(SDTrainerExtension):
|
||||
uid = "textual_inversion_trainer"
|
||||
@@ -26,5 +43,5 @@ class TextualInversionTrainer(SDTrainerExtension):
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
SDTrainerExtension, TextualInversionTrainer
|
||||
SDTrainerExtension, TextualInversionTrainer, MultiGPUSDTrainerExtension
|
||||
]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
||||
import gc
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from typing import ForwardRef, List
|
||||
from typing import ForwardRef, List, Optional, Union
|
||||
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
@@ -22,6 +22,7 @@ class GenerateConfig:
|
||||
self.sampler = kwargs.get('sampler', 'ddpm')
|
||||
self.width = kwargs.get('width', 512)
|
||||
self.height = kwargs.get('height', 512)
|
||||
self.size_list: Union[List[int], None] = kwargs.get('size_list', None)
|
||||
self.neg = kwargs.get('neg', '')
|
||||
self.seed = kwargs.get('seed', -1)
|
||||
self.guidance_scale = kwargs.get('guidance_scale', 7)
|
||||
@@ -30,18 +31,33 @@ class GenerateConfig:
|
||||
self.neg_2 = kwargs.get('neg_2', None)
|
||||
self.prompts = kwargs.get('prompts', None)
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.compile = kwargs.get('compile', False)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.prompt_file = kwargs.get('prompt_file', False)
|
||||
self.prompts_in_file = self.prompts
|
||||
if self.prompts is None:
|
||||
raise ValueError("Prompts must be set")
|
||||
if isinstance(self.prompts, str):
|
||||
if os.path.exists(self.prompts):
|
||||
with open(self.prompts, 'r', encoding='utf-8') as f:
|
||||
self.prompts = f.read().splitlines()
|
||||
self.prompts = [p.strip() for p in self.prompts if len(p.strip()) > 0]
|
||||
self.prompts_in_file = f.read().splitlines()
|
||||
self.prompts_in_file = [p.strip() for p in self.prompts_in_file if len(p.strip()) > 0]
|
||||
else:
|
||||
raise ValueError("Prompts file does not exist, put in list if you want to use a list of prompts")
|
||||
|
||||
self.random_prompts = kwargs.get('random_prompts', False)
|
||||
self.max_random_per_prompt = kwargs.get('max_random_per_prompt', 1)
|
||||
self.max_images = kwargs.get('max_images', 10000)
|
||||
|
||||
if self.random_prompts:
|
||||
self.prompts = []
|
||||
for i in range(self.max_images):
|
||||
num_prompts = random.randint(1, self.max_random_per_prompt)
|
||||
prompt_list = [random.choice(self.prompts_in_file) for _ in range(num_prompts)]
|
||||
self.prompts.append(", ".join(prompt_list))
|
||||
else:
|
||||
self.prompts = self.prompts_in_file
|
||||
|
||||
if kwargs.get('shuffle', False):
|
||||
# shuffle the prompts
|
||||
random.shuffle(self.prompts)
|
||||
@@ -64,6 +80,7 @@ class GenerateProcess(BaseProcess):
|
||||
self.model_config = ModelConfig(**self.get_conf('model', required=True))
|
||||
self.device = self.get_conf('device', self.job.device)
|
||||
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
|
||||
self.torch_dtype = get_torch_dtype(self.get_conf('dtype', 'float16'))
|
||||
|
||||
self.progress_bar = None
|
||||
self.sd = StableDiffusion(
|
||||
@@ -71,37 +88,57 @@ class GenerateProcess(BaseProcess):
|
||||
model_config=self.model_config,
|
||||
dtype=self.model_config.dtype,
|
||||
)
|
||||
|
||||
print(f"Using device {self.device}")
|
||||
|
||||
def clean_prompt(self, prompt: str):
|
||||
# remove any non alpha numeric characters or ,'" from prompt
|
||||
return ''.join(e for e in prompt if e.isalnum() or e in ", '\"")
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
with torch.no_grad():
|
||||
super().run()
|
||||
print("Loading model...")
|
||||
self.sd.load_model()
|
||||
self.sd.pipeline.to(self.device, self.torch_dtype)
|
||||
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for prompt in self.generate_config.prompts:
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=self.generate_config.width,
|
||||
height=self.generate_config.height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
|
||||
print("Compiling model...")
|
||||
# self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead", fullgraph=True)
|
||||
if self.generate_config.compile:
|
||||
self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead")
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
print(f"Generating {len(self.generate_config.prompts)} images")
|
||||
# build prompt image configs
|
||||
prompt_image_configs = []
|
||||
for prompt in self.generate_config.prompts:
|
||||
width = self.generate_config.width
|
||||
height = self.generate_config.height
|
||||
prompt = self.clean_prompt(prompt)
|
||||
|
||||
if self.generate_config.size_list is not None:
|
||||
# randomly select a size
|
||||
width, height = random.choice(self.generate_config.size_list)
|
||||
|
||||
prompt_image_configs.append(GenerateImageConfig(
|
||||
prompt=prompt,
|
||||
prompt_2=self.generate_config.prompt_2,
|
||||
width=width,
|
||||
height=height,
|
||||
num_inference_steps=self.generate_config.sample_steps,
|
||||
guidance_scale=self.generate_config.guidance_scale,
|
||||
negative_prompt=self.generate_config.neg,
|
||||
negative_prompt_2=self.generate_config.neg_2,
|
||||
seed=self.generate_config.seed,
|
||||
guidance_rescale=self.generate_config.guidance_rescale,
|
||||
output_ext=self.generate_config.ext,
|
||||
output_folder=self.output_folder,
|
||||
add_prompt_file=self.generate_config.prompt_file
|
||||
))
|
||||
# generate images
|
||||
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
|
||||
|
||||
print("Done generating images")
|
||||
# cleanup
|
||||
del self.sd
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -371,7 +371,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
|
||||
# ger a random number of steps
|
||||
timesteps_to = torch.randint(
|
||||
1, self.train_config.max_denoising_steps, (1,)
|
||||
1, self.train_config.max_denoising_steps - 1, (1,)
|
||||
).item()
|
||||
|
||||
# get noise
|
||||
@@ -389,7 +389,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
assert not self.network.is_active
|
||||
self.sd.unet.eval()
|
||||
# pass the multiplier list to the network
|
||||
self.network.multiplier = prompt_pair.multiplier_list
|
||||
# double up since we are doing cfg
|
||||
self.network.multiplier = prompt_pair.multiplier_list + prompt_pair.multiplier_list
|
||||
denoised_latents = self.sd.diffuse_some_steps(
|
||||
latents, # pass simple noise latents
|
||||
train_tools.concat_prompt_embeddings(
|
||||
@@ -507,7 +508,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
for anchor_chunk, denoised_latent_chunk, anchor_target_noise_chunk in zip(
|
||||
anchor_chunks, denoised_latent_chunks, anchor_target_noise_chunks
|
||||
):
|
||||
self.network.multiplier = anchor_chunk.multiplier_list
|
||||
self.network.multiplier = anchor_chunk.multiplier_list + anchor_chunk.multiplier_list
|
||||
|
||||
anchor_pred_noise = get_noise_pred(
|
||||
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
|
||||
@@ -582,7 +583,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
mask_multiplier_chunks,
|
||||
unmasked_target_chunks
|
||||
):
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list
|
||||
self.network.multiplier = prompt_pair_chunk.multiplier_list + prompt_pair_chunk.multiplier_list
|
||||
target_latents = get_noise_pred(
|
||||
prompt_pair_chunk.positive_target,
|
||||
prompt_pair_chunk.target_class,
|
||||
@@ -611,6 +612,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
offset_neutral = neutral_latents_chunk
|
||||
# offsets are already adjusted on a per-batch basis
|
||||
offset_neutral += offset
|
||||
offset_neutral = offset_neutral.detach().requires_grad_(False)
|
||||
|
||||
# 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
|
||||
loss = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import copy
|
||||
import glob
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
|
||||
@@ -13,6 +14,7 @@ from torch import nn
|
||||
from torchvision.transforms import transforms
|
||||
|
||||
from jobs.process import BaseTrainProcess
|
||||
from toolkit.image_utils import show_tensors
|
||||
from toolkit.kohya_model_util import load_vae, convert_diffusers_back_to_ldm
|
||||
from toolkit.data_loader import ImageDataset
|
||||
from toolkit.losses import ComparativeTotalVariation, get_gradient_penalty, PatternLoss
|
||||
@@ -25,6 +27,8 @@ from tqdm import tqdm
|
||||
import time
|
||||
import numpy as np
|
||||
from .models.vgg19_critic import Critic
|
||||
from torchvision.transforms import Resize
|
||||
import lpips
|
||||
|
||||
IMAGE_TRANSFORMS = transforms.Compose(
|
||||
[
|
||||
@@ -62,6 +66,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.kld_weight = self.get_conf('kld_weight', 0, as_type=float)
|
||||
self.mse_weight = self.get_conf('mse_weight', 1e0, as_type=float)
|
||||
self.tv_weight = self.get_conf('tv_weight', 1e0, as_type=float)
|
||||
self.lpips_weight = self.get_conf('lpips_weight', 1e0, as_type=float)
|
||||
self.critic_weight = self.get_conf('critic_weight', 1, as_type=float)
|
||||
self.pattern_weight = self.get_conf('pattern_weight', 1, as_type=float)
|
||||
self.optimizer_params = self.get_conf('optimizer_params', {})
|
||||
@@ -71,6 +76,9 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.vgg_19 = None
|
||||
self.style_weight_scalers = []
|
||||
self.content_weight_scalers = []
|
||||
self.lpips_loss:lpips.LPIPS = None
|
||||
|
||||
self.vae_scale_factor = 8
|
||||
|
||||
self.step_num = 0
|
||||
self.epoch_num = 0
|
||||
@@ -137,6 +145,15 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
num_workers=6
|
||||
)
|
||||
|
||||
def remove_oldest_checkpoint(self):
|
||||
max_to_keep = 4
|
||||
folders = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
|
||||
if len(folders) > max_to_keep:
|
||||
folders.sort(key=os.path.getmtime)
|
||||
for folder in folders[:-max_to_keep]:
|
||||
print(f"Removing {folder}")
|
||||
shutil.rmtree(folder)
|
||||
|
||||
def setup_vgg19(self):
|
||||
if self.vgg_19 is None:
|
||||
self.vgg_19, self.style_losses, self.content_losses, self.vgg19_pool_4 = get_style_model_and_losses(
|
||||
@@ -211,7 +228,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
|
||||
def get_pattern_loss(self, pred, target):
|
||||
if self._pattern_loss is None:
|
||||
self._pattern_loss = PatternLoss(pattern_size=8, dtype=self.torch_dtype).to(self.device,
|
||||
self._pattern_loss = PatternLoss(pattern_size=16, dtype=self.torch_dtype).to(self.device,
|
||||
dtype=self.torch_dtype)
|
||||
loss = torch.mean(self._pattern_loss(pred, target))
|
||||
return loss
|
||||
@@ -226,25 +243,21 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
step_num = f"_{str(step).zfill(9)}"
|
||||
|
||||
self.update_training_metadata()
|
||||
filename = f'{self.job.name}{step_num}.safetensors'
|
||||
# prepare meta
|
||||
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
|
||||
filename = f'{self.job.name}{step_num}_diffusers'
|
||||
|
||||
state_dict = convert_diffusers_back_to_ldm(self.vae)
|
||||
|
||||
for key in list(state_dict.keys()):
|
||||
v = state_dict[key]
|
||||
v = v.detach().clone().to("cpu").to(torch.float32)
|
||||
state_dict[key] = v
|
||||
|
||||
# having issues with meta
|
||||
save_file(state_dict, os.path.join(self.save_root, filename), save_meta)
|
||||
self.vae = self.vae.to("cpu", dtype=torch.float16)
|
||||
self.vae.save_pretrained(
|
||||
save_directory=os.path.join(self.save_root, filename)
|
||||
)
|
||||
self.vae = self.vae.to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
self.print(f"Saved to {os.path.join(self.save_root, filename)}")
|
||||
|
||||
if self.use_critic:
|
||||
self.critic.save(step)
|
||||
|
||||
self.remove_oldest_checkpoint()
|
||||
|
||||
def sample(self, step=None):
|
||||
sample_folder = os.path.join(self.save_root, 'samples')
|
||||
if not os.path.exists(sample_folder):
|
||||
@@ -280,6 +293,13 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
output_img.paste(input_img, (0, 0))
|
||||
output_img.paste(decoded, (self.resolution, 0))
|
||||
|
||||
scale_up = 2
|
||||
if output_img.height <= 300:
|
||||
scale_up = 4
|
||||
|
||||
# scale up using nearest neighbor
|
||||
output_img = output_img.resize((output_img.width * scale_up, output_img.height * scale_up), Image.NEAREST)
|
||||
|
||||
step_num = ''
|
||||
if step is not None:
|
||||
# zero-pad 9 digits
|
||||
@@ -294,7 +314,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
path_to_load = self.vae_path
|
||||
# see if we have a checkpoint in out output to resume from
|
||||
self.print(f"Looking for latest checkpoint in {self.save_root}")
|
||||
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*.safetensors"))
|
||||
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
|
||||
if files and len(files) > 0:
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
print(f" - Latest checkpoint is: {latest_file}")
|
||||
@@ -306,13 +326,14 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.print(f"Loading VAE")
|
||||
self.print(f" - Loading VAE: {path_to_load}")
|
||||
if self.vae is None:
|
||||
self.vae = load_vae(path_to_load, dtype=self.torch_dtype)
|
||||
self.vae = AutoencoderKL.from_pretrained(path_to_load)
|
||||
|
||||
# set decoder to train
|
||||
self.vae.to(self.device, dtype=self.torch_dtype)
|
||||
self.vae.requires_grad_(False)
|
||||
self.vae.eval()
|
||||
self.vae.decoder.train()
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config['block_out_channels']) - 1)
|
||||
|
||||
def run(self):
|
||||
super().run()
|
||||
@@ -374,6 +395,10 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
if self.use_critic:
|
||||
self.critic.setup()
|
||||
|
||||
if self.lpips_weight > 0 and self.lpips_loss is None:
|
||||
# self.lpips_loss = lpips.LPIPS(net='vgg')
|
||||
self.lpips_loss = lpips.LPIPS(net='vgg').to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
|
||||
optimizer_params=self.optimizer_params)
|
||||
|
||||
@@ -397,6 +422,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.sample()
|
||||
blank_losses = OrderedDict({
|
||||
"total": [],
|
||||
"lpips": [],
|
||||
"style": [],
|
||||
"content": [],
|
||||
"mse": [],
|
||||
@@ -415,17 +441,29 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
for batch in self.data_loader:
|
||||
if self.step_num >= self.max_steps:
|
||||
break
|
||||
with torch.no_grad():
|
||||
|
||||
batch = batch.to(self.device, dtype=self.torch_dtype)
|
||||
batch = batch.to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
# forward pass
|
||||
dgd = self.vae.encode(batch).latent_dist
|
||||
mu, logvar = dgd.mean, dgd.logvar
|
||||
latents = dgd.sample()
|
||||
latents.requires_grad_(True)
|
||||
# resize so it matches size of vae evenly
|
||||
if batch.shape[2] % self.vae_scale_factor != 0 or batch.shape[3] % self.vae_scale_factor != 0:
|
||||
batch = Resize((batch.shape[2] // self.vae_scale_factor * self.vae_scale_factor,
|
||||
batch.shape[3] // self.vae_scale_factor * self.vae_scale_factor))(batch)
|
||||
|
||||
# forward pass
|
||||
dgd = self.vae.encode(batch).latent_dist
|
||||
mu, logvar = dgd.mean, dgd.logvar
|
||||
latents = dgd.sample()
|
||||
latents.detach().requires_grad_(True)
|
||||
|
||||
pred = self.vae.decode(latents).sample
|
||||
|
||||
with torch.no_grad():
|
||||
show_tensors(
|
||||
pred.clamp(-1, 1).clone(),
|
||||
"combined tensor"
|
||||
)
|
||||
|
||||
# Run through VGG19
|
||||
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
|
||||
stacked = torch.cat([pred, batch], dim=0)
|
||||
@@ -441,14 +479,31 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
content_loss = self.get_content_loss() * self.content_weight
|
||||
kld_loss = self.get_kld_loss(mu, logvar) * self.kld_weight
|
||||
mse_loss = self.get_mse_loss(pred, batch) * self.mse_weight
|
||||
if self.lpips_weight > 0:
|
||||
lpips_loss = self.lpips_loss(
|
||||
pred.clamp(-1, 1),
|
||||
batch.clamp(-1, 1)
|
||||
).mean() * self.lpips_weight
|
||||
else:
|
||||
lpips_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
|
||||
tv_loss = self.get_tv_loss(pred, batch) * self.tv_weight
|
||||
pattern_loss = self.get_pattern_loss(pred, batch) * self.pattern_weight
|
||||
if self.use_critic:
|
||||
critic_gen_loss = self.critic.get_critic_loss(self.vgg19_pool_4.tensor) * self.critic_weight
|
||||
|
||||
# do not let abs critic gen loss be higher than abs lpips * 0.1 if using it
|
||||
if self.lpips_weight > 0:
|
||||
max_target = lpips_loss.abs() * 0.1
|
||||
with torch.no_grad():
|
||||
crit_g_scaler = 1.0
|
||||
if critic_gen_loss.abs() > max_target:
|
||||
crit_g_scaler = max_target / critic_gen_loss.abs()
|
||||
|
||||
critic_gen_loss *= crit_g_scaler
|
||||
else:
|
||||
critic_gen_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
|
||||
|
||||
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss
|
||||
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss + lpips_loss
|
||||
|
||||
# Backward pass and optimization
|
||||
optimizer.zero_grad()
|
||||
@@ -460,6 +515,8 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
loss_value = loss.item()
|
||||
# get exponent like 3.54e-4
|
||||
loss_string = f"loss: {loss_value:.2e}"
|
||||
if self.lpips_weight > 0:
|
||||
loss_string += f" lpips: {lpips_loss.item():.2e}"
|
||||
if self.content_weight > 0:
|
||||
loss_string += f" cnt: {content_loss.item():.2e}"
|
||||
if self.style_weight > 0:
|
||||
@@ -477,7 +534,8 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
if self.use_critic:
|
||||
loss_string += f" crD: {critic_d_loss:.2e}"
|
||||
|
||||
if self.optimizer_type.startswith('dadaptation'):
|
||||
if self.optimizer_type.startswith('dadaptation') or \
|
||||
self.optimizer_type.lower().startswith('prodigy'):
|
||||
learning_rate = (
|
||||
optimizer.param_groups[0]["d"] *
|
||||
optimizer.param_groups[0]["lr"]
|
||||
@@ -495,6 +553,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
self.progress_bar.update(1)
|
||||
|
||||
epoch_losses["total"].append(loss_value)
|
||||
epoch_losses["lpips"].append(lpips_loss.item())
|
||||
epoch_losses["style"].append(style_loss.item())
|
||||
epoch_losses["content"].append(content_loss.item())
|
||||
epoch_losses["mse"].append(mse_loss.item())
|
||||
@@ -505,6 +564,7 @@ class TrainVAEProcess(BaseTrainProcess):
|
||||
epoch_losses["crD"].append(critic_d_loss)
|
||||
|
||||
log_losses["total"].append(loss_value)
|
||||
log_losses["lpips"].append(lpips_loss.item())
|
||||
log_losses["style"].append(style_loss.item())
|
||||
log_losses["content"].append(content_loss.item())
|
||||
log_losses["mse"].append(mse_loss.item())
|
||||
|
||||
291
notebooks/FLUX_1_dev_LoRA_Training.ipynb
Normal file
291
notebooks/FLUX_1_dev_LoRA_Training.ipynb
Normal file
@@ -0,0 +1,291 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"id": "zl-S0m3pkQC5"
|
||||
},
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## FLUX.1-dev Training\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit\n",
|
||||
"!mkdir -p /content/dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "UFUW4ZMmnp1V"
|
||||
},
|
||||
"source": [
|
||||
"Put your image dataset in the `/content/dataset` folder"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "OV0HnOI6o8V6"
|
||||
},
|
||||
"source": [
|
||||
"## Model License\n",
|
||||
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
|
||||
"\n",
|
||||
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
|
||||
"\n",
|
||||
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3yZZdhFRoj2m"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import getpass\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"# Prompt for the token\n",
|
||||
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
|
||||
"\n",
|
||||
"# Set the environment variable\n",
|
||||
"os.environ['HF_TOKEN'] = hf_token\n",
|
||||
"\n",
|
||||
"print(\"HF_TOKEN environment variable has been set.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "9gO2EzQ1kQC8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image\n",
|
||||
"import os\n",
|
||||
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
},
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict([\n",
|
||||
" ('job', 'extension'),\n",
|
||||
" ('config', OrderedDict([\n",
|
||||
" # this name will be the folder and filename name\n",
|
||||
" ('name', 'my_first_flux_lora_v1'),\n",
|
||||
" ('process', [\n",
|
||||
" OrderedDict([\n",
|
||||
" ('type', 'sd_trainer'),\n",
|
||||
" # root folder to save training sessions/samples/weights\n",
|
||||
" ('training_folder', '/content/output'),\n",
|
||||
" # uncomment to see performance stats in the terminal every N steps\n",
|
||||
" #('performance_log_every', 1000),\n",
|
||||
" ('device', 'cuda:0'),\n",
|
||||
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
|
||||
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
|
||||
" # ('trigger_word', 'image'),\n",
|
||||
" ('network', OrderedDict([\n",
|
||||
" ('type', 'lora'),\n",
|
||||
" ('linear', 16),\n",
|
||||
" ('linear_alpha', 16)\n",
|
||||
" ])),\n",
|
||||
" ('save', OrderedDict([\n",
|
||||
" ('dtype', 'float16'), # precision to save\n",
|
||||
" ('save_every', 250), # save every this many steps\n",
|
||||
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
|
||||
" ])),\n",
|
||||
" ('datasets', [\n",
|
||||
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
|
||||
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
|
||||
" # images will automatically be resized and bucketed into the resolution specified\n",
|
||||
" OrderedDict([\n",
|
||||
" ('folder_path', '/content/dataset'),\n",
|
||||
" ('caption_ext', 'txt'),\n",
|
||||
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
|
||||
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
|
||||
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
|
||||
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
|
||||
" ])\n",
|
||||
" ]),\n",
|
||||
" ('train', OrderedDict([\n",
|
||||
" ('batch_size', 1),\n",
|
||||
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
|
||||
" ('gradient_accumulation_steps', 1),\n",
|
||||
" ('train_unet', True),\n",
|
||||
" ('train_text_encoder', False), # probably won't work with flux\n",
|
||||
" ('content_or_style', 'balanced'), # content, style, balanced\n",
|
||||
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
|
||||
" ('noise_scheduler', 'flowmatch'), # for training only\n",
|
||||
" ('optimizer', 'adamw8bit'),\n",
|
||||
" ('lr', 1e-4),\n",
|
||||
"\n",
|
||||
" # uncomment this to skip the pre training sample\n",
|
||||
" # ('skip_first_sample', True),\n",
|
||||
"\n",
|
||||
" # uncomment to completely disable sampling\n",
|
||||
" # ('disable_sampling', True),\n",
|
||||
"\n",
|
||||
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
|
||||
" # ('linear_timesteps', True),\n",
|
||||
"\n",
|
||||
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
|
||||
" ('ema_config', OrderedDict([\n",
|
||||
" ('use_ema', True),\n",
|
||||
" ('ema_decay', 0.99)\n",
|
||||
" ])),\n",
|
||||
"\n",
|
||||
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
|
||||
" ('dtype', 'bf16')\n",
|
||||
" ])),\n",
|
||||
" ('model', OrderedDict([\n",
|
||||
" # huggingface model name or path\n",
|
||||
" ('name_or_path', 'black-forest-labs/FLUX.1-dev'),\n",
|
||||
" ('is_flux', True),\n",
|
||||
" ('quantize', True), # run 8bit mixed precision\n",
|
||||
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
|
||||
" ])),\n",
|
||||
" ('sample', OrderedDict([\n",
|
||||
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
|
||||
" ('sample_every', 250), # sample every this many steps\n",
|
||||
" ('width', 1024),\n",
|
||||
" ('height', 1024),\n",
|
||||
" ('prompts', [\n",
|
||||
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
|
||||
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
|
||||
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
|
||||
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
|
||||
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
|
||||
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
|
||||
" 'a bear building a log cabin in the snow covered mountains',\n",
|
||||
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
|
||||
" 'hipster man with a beard, building a chair, in a wood shop',\n",
|
||||
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
|
||||
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
|
||||
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
|
||||
" ]),\n",
|
||||
" ('neg', ''), # not used on flux\n",
|
||||
" ('seed', 42),\n",
|
||||
" ('walk_seed', True),\n",
|
||||
" ('guidance_scale', 4),\n",
|
||||
" ('sample_steps', 20)\n",
|
||||
" ]))\n",
|
||||
" ])\n",
|
||||
" ])\n",
|
||||
" ])),\n",
|
||||
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
|
||||
" ('meta', OrderedDict([\n",
|
||||
" ('name', '[name]'),\n",
|
||||
" ('version', '1.0')\n",
|
||||
" ]))\n",
|
||||
"])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
},
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
},
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"accelerator": "GPU",
|
||||
"colab": {
|
||||
"gpuType": "A100",
|
||||
"machine_shape": "hm",
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
296
notebooks/FLUX_1_schnell_LoRA_Training.ipynb
Normal file
296
notebooks/FLUX_1_schnell_LoRA_Training.ipynb
Normal file
@@ -0,0 +1,296 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"collapsed": false,
|
||||
"id": "zl-S0m3pkQC5"
|
||||
},
|
||||
"source": [
|
||||
"# AI Toolkit by Ostris\n",
|
||||
"## FLUX.1-schnell Training\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3cokMT-WC6rG"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!nvidia-smi"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"id": "BvAG0GKAh59G"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!git clone https://github.com/ostris/ai-toolkit\n",
|
||||
"!mkdir -p /content/dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "UFUW4ZMmnp1V"
|
||||
},
|
||||
"source": [
|
||||
"Put your image dataset in the `/content/dataset` folder"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"collapsed": true,
|
||||
"id": "XGZqVER_aQJW"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "OV0HnOI6o8V6"
|
||||
},
|
||||
"source": [
|
||||
"## Model License\n",
|
||||
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
|
||||
"\n",
|
||||
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
|
||||
"\n",
|
||||
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "3yZZdhFRoj2m"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import getpass\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"# Prompt for the token\n",
|
||||
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
|
||||
"\n",
|
||||
"# Set the environment variable\n",
|
||||
"os.environ['HF_TOKEN'] = hf_token\n",
|
||||
"\n",
|
||||
"print(\"HF_TOKEN environment variable has been set.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 5,
|
||||
"metadata": {
|
||||
"id": "9gO2EzQ1kQC8"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"import sys\n",
|
||||
"sys.path.append('/content/ai-toolkit')\n",
|
||||
"from toolkit.job import run_job\n",
|
||||
"from collections import OrderedDict\n",
|
||||
"from PIL import Image\n",
|
||||
"import os\n",
|
||||
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "N8UUFzVRigbC"
|
||||
},
|
||||
"source": [
|
||||
"## Setup\n",
|
||||
"\n",
|
||||
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 6,
|
||||
"metadata": {
|
||||
"id": "_t28QURYjRQO"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from collections import OrderedDict\n",
|
||||
"\n",
|
||||
"job_to_run = OrderedDict([\n",
|
||||
" ('job', 'extension'),\n",
|
||||
" ('config', OrderedDict([\n",
|
||||
" # this name will be the folder and filename name\n",
|
||||
" ('name', 'my_first_flux_lora_v1'),\n",
|
||||
" ('process', [\n",
|
||||
" OrderedDict([\n",
|
||||
" ('type', 'sd_trainer'),\n",
|
||||
" # root folder to save training sessions/samples/weights\n",
|
||||
" ('training_folder', '/content/output'),\n",
|
||||
" # uncomment to see performance stats in the terminal every N steps\n",
|
||||
" #('performance_log_every', 1000),\n",
|
||||
" ('device', 'cuda:0'),\n",
|
||||
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
|
||||
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
|
||||
" # ('trigger_word', 'image'),\n",
|
||||
" ('network', OrderedDict([\n",
|
||||
" ('type', 'lora'),\n",
|
||||
" ('linear', 16),\n",
|
||||
" ('linear_alpha', 16)\n",
|
||||
" ])),\n",
|
||||
" ('save', OrderedDict([\n",
|
||||
" ('dtype', 'float16'), # precision to save\n",
|
||||
" ('save_every', 250), # save every this many steps\n",
|
||||
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
|
||||
" ])),\n",
|
||||
" ('datasets', [\n",
|
||||
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
|
||||
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
|
||||
" # images will automatically be resized and bucketed into the resolution specified\n",
|
||||
" OrderedDict([\n",
|
||||
" ('folder_path', '/content/dataset'),\n",
|
||||
" ('caption_ext', 'txt'),\n",
|
||||
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
|
||||
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
|
||||
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
|
||||
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
|
||||
" ])\n",
|
||||
" ]),\n",
|
||||
" ('train', OrderedDict([\n",
|
||||
" ('batch_size', 1),\n",
|
||||
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
|
||||
" ('gradient_accumulation_steps', 1),\n",
|
||||
" ('train_unet', True),\n",
|
||||
" ('train_text_encoder', False), # probably won't work with flux\n",
|
||||
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
|
||||
" ('noise_scheduler', 'flowmatch'), # for training only\n",
|
||||
" ('optimizer', 'adamw8bit'),\n",
|
||||
" ('lr', 1e-4),\n",
|
||||
"\n",
|
||||
" # uncomment this to skip the pre training sample\n",
|
||||
" # ('skip_first_sample', True),\n",
|
||||
"\n",
|
||||
" # uncomment to completely disable sampling\n",
|
||||
" # ('disable_sampling', True),\n",
|
||||
"\n",
|
||||
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
|
||||
" # ('linear_timesteps', True),\n",
|
||||
"\n",
|
||||
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
|
||||
" ('ema_config', OrderedDict([\n",
|
||||
" ('use_ema', True),\n",
|
||||
" ('ema_decay', 0.99)\n",
|
||||
" ])),\n",
|
||||
"\n",
|
||||
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
|
||||
" ('dtype', 'bf16')\n",
|
||||
" ])),\n",
|
||||
" ('model', OrderedDict([\n",
|
||||
" # huggingface model name or path\n",
|
||||
" ('name_or_path', 'black-forest-labs/FLUX.1-schnell'),\n",
|
||||
" ('assistant_lora_path', 'ostris/FLUX.1-schnell-training-adapter'), # Required for flux schnell training\n",
|
||||
" ('is_flux', True),\n",
|
||||
" ('quantize', True), # run 8bit mixed precision\n",
|
||||
" # low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary\n",
|
||||
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
|
||||
" ])),\n",
|
||||
" ('sample', OrderedDict([\n",
|
||||
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
|
||||
" ('sample_every', 250), # sample every this many steps\n",
|
||||
" ('width', 1024),\n",
|
||||
" ('height', 1024),\n",
|
||||
" ('prompts', [\n",
|
||||
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
|
||||
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
|
||||
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
|
||||
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
|
||||
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
|
||||
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
|
||||
" 'a bear building a log cabin in the snow covered mountains',\n",
|
||||
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
|
||||
" 'hipster man with a beard, building a chair, in a wood shop',\n",
|
||||
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
|
||||
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
|
||||
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
|
||||
" ]),\n",
|
||||
" ('neg', ''), # not used on flux\n",
|
||||
" ('seed', 42),\n",
|
||||
" ('walk_seed', True),\n",
|
||||
" ('guidance_scale', 1), # schnell does not do guidance\n",
|
||||
" ('sample_steps', 4) # 1 - 4 works well\n",
|
||||
" ]))\n",
|
||||
" ])\n",
|
||||
" ])\n",
|
||||
" ])),\n",
|
||||
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
|
||||
" ('meta', OrderedDict([\n",
|
||||
" ('name', '[name]'),\n",
|
||||
" ('version', '1.0')\n",
|
||||
" ]))\n",
|
||||
"])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "h6F1FlM2Wb3l"
|
||||
},
|
||||
"source": [
|
||||
"## Run it\n",
|
||||
"\n",
|
||||
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {
|
||||
"id": "HkajwI8gteOh"
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"run_job(job_to_run)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
"id": "Hblgb5uwW5SD"
|
||||
},
|
||||
"source": [
|
||||
"## Done\n",
|
||||
"\n",
|
||||
"Check your ourput dir and get your slider\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"accelerator": "GPU",
|
||||
"colab": {
|
||||
"gpuType": "A100",
|
||||
"machine_shape": "hm",
|
||||
"provenance": []
|
||||
},
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 0
|
||||
}
|
||||
Submodule repositories/ipadapter updated: d8ab37c421...5a18b1f366
@@ -1,8 +1,8 @@
|
||||
torch
|
||||
torchvision
|
||||
safetensors
|
||||
diffusers==0.21.3
|
||||
git+https://github.com/huggingface/transformers.git
|
||||
git+https://github.com/huggingface/diffusers.git
|
||||
transformers
|
||||
lycoris-lora==1.8.3
|
||||
flatten_json
|
||||
pyyaml
|
||||
@@ -21,4 +21,12 @@ open_clip_torch
|
||||
timm
|
||||
prodigyopt
|
||||
controlnet_aux==0.0.7
|
||||
python-dotenv
|
||||
python-dotenv
|
||||
bitsandbytes
|
||||
hf_transfer
|
||||
lpips
|
||||
pytorch_fid
|
||||
optimum-quanto
|
||||
sentencepiece
|
||||
huggingface_hub
|
||||
peft
|
||||
1
run.py
1
run.py
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
from typing import Union, OrderedDict
|
||||
from dotenv import load_dotenv
|
||||
|
||||
175
run_modal.py
Normal file
175
run_modal.py
Normal file
@@ -0,0 +1,175 @@
|
||||
'''
|
||||
|
||||
ostris/ai-toolkit on https://modal.com
|
||||
Run training with the following command:
|
||||
modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml
|
||||
|
||||
'''
|
||||
|
||||
import os
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
import modal
|
||||
from dotenv import load_dotenv
|
||||
# Load the .env file if it exists
|
||||
load_dotenv()
|
||||
|
||||
sys.path.insert(0, "/root/ai-toolkit")
|
||||
# must come before ANY torch or fastai imports
|
||||
# import toolkit.cuda_malloc
|
||||
|
||||
# turn off diffusers telemetry until I can figure out how to make it opt-in
|
||||
os.environ['DISABLE_TELEMETRY'] = 'YES'
|
||||
|
||||
# define the volume for storing model outputs, using "creating volumes lazily": https://modal.com/docs/guide/volumes
|
||||
# you will find your model, samples and optimizer stored in: https://modal.com/storage/your-username/main/flux-lora-models
|
||||
model_volume = modal.Volume.from_name("flux-lora-models", create_if_missing=True)
|
||||
|
||||
# modal_output, due to "cannot mount volume on non-empty path" requirement
|
||||
MOUNT_DIR = "/root/ai-toolkit/modal_output" # modal_output, due to "cannot mount volume on non-empty path" requirement
|
||||
|
||||
# define modal app
|
||||
image = (
|
||||
modal.Image.debian_slim(python_version="3.11")
|
||||
# install required system and pip packages, more about this modal approach: https://modal.com/docs/examples/dreambooth_app
|
||||
.apt_install("libgl1", "libglib2.0-0")
|
||||
.pip_install(
|
||||
"python-dotenv",
|
||||
"torch",
|
||||
"diffusers[torch]",
|
||||
"transformers",
|
||||
"ftfy",
|
||||
"torchvision",
|
||||
"oyaml",
|
||||
"opencv-python",
|
||||
"albumentations",
|
||||
"safetensors",
|
||||
"lycoris-lora==1.8.3",
|
||||
"flatten_json",
|
||||
"pyyaml",
|
||||
"tensorboard",
|
||||
"kornia",
|
||||
"invisible-watermark",
|
||||
"einops",
|
||||
"accelerate",
|
||||
"toml",
|
||||
"pydantic",
|
||||
"omegaconf",
|
||||
"k-diffusion",
|
||||
"open_clip_torch",
|
||||
"timm",
|
||||
"prodigyopt",
|
||||
"controlnet_aux==0.0.7",
|
||||
"bitsandbytes",
|
||||
"hf_transfer",
|
||||
"lpips",
|
||||
"pytorch_fid",
|
||||
"optimum-quanto",
|
||||
"sentencepiece",
|
||||
"huggingface_hub",
|
||||
"peft"
|
||||
)
|
||||
)
|
||||
|
||||
# mount for the entire ai-toolkit directory
|
||||
# example: "/Users/username/ai-toolkit" is the local directory, "/root/ai-toolkit" is the remote directory
|
||||
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
|
||||
|
||||
# create the Modal app with the necessary mounts and volumes
|
||||
app = modal.App(name="flux-lora-training", image=image, mounts=[code_mount], volumes={MOUNT_DIR: model_volume})
|
||||
|
||||
# Check if we have DEBUG_TOOLKIT in env
|
||||
if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
# Set torch to trace mode
|
||||
import torch
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
|
||||
def print_end_message(jobs_completed, jobs_failed):
|
||||
failure_string = f"{jobs_failed} failure{'' if jobs_failed == 1 else 's'}" if jobs_failed > 0 else ""
|
||||
completed_string = f"{jobs_completed} completed job{'' if jobs_completed == 1 else 's'}"
|
||||
|
||||
print("")
|
||||
print("========================================")
|
||||
print("Result:")
|
||||
if len(completed_string) > 0:
|
||||
print(f" - {completed_string}")
|
||||
if len(failure_string) > 0:
|
||||
print(f" - {failure_string}")
|
||||
print("========================================")
|
||||
|
||||
|
||||
@app.function(
|
||||
# request a GPU with at least 24GB VRAM
|
||||
# more about modal GPU's: https://modal.com/docs/guide/gpu
|
||||
gpu="A100", # gpu="H100"
|
||||
# more about modal timeouts: https://modal.com/docs/guide/timeouts
|
||||
timeout=7200 # 2 hours, increase or decrease if needed
|
||||
)
|
||||
def main(config_file_list_str: str, recover: bool = False, name: str = None):
|
||||
# convert the config file list from a string to a list
|
||||
config_file_list = config_file_list_str.split(",")
|
||||
|
||||
jobs_completed = 0
|
||||
jobs_failed = 0
|
||||
|
||||
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
job = get_job(config_file, name)
|
||||
|
||||
job.config['process'][0]['training_folder'] = MOUNT_DIR
|
||||
os.makedirs(MOUNT_DIR, exist_ok=True)
|
||||
print(f"Training outputs will be saved to: {MOUNT_DIR}")
|
||||
|
||||
# run the job
|
||||
job.run()
|
||||
|
||||
# commit the volume after training
|
||||
model_volume.commit()
|
||||
|
||||
job.cleanup()
|
||||
jobs_completed += 1
|
||||
except Exception as e:
|
||||
print(f"Error running job: {e}")
|
||||
jobs_failed += 1
|
||||
if not recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# require at least one config file
|
||||
parser.add_argument(
|
||||
'config_file_list',
|
||||
nargs='+',
|
||||
type=str,
|
||||
help='Name of config file (eg: person_v1 for config/person_v1.json/yaml), or full path if it is not in config folder, you can pass multiple config files and run them all sequentially'
|
||||
)
|
||||
|
||||
# flag to continue if a job fails
|
||||
parser.add_argument(
|
||||
'-r', '--recover',
|
||||
action='store_true',
|
||||
help='Continue running additional jobs even if a job fails'
|
||||
)
|
||||
|
||||
# optional name replacement for config file
|
||||
parser.add_argument(
|
||||
'-n', '--name',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# convert list of config files to a comma-separated string for Modal compatibility
|
||||
config_file_list_str = ",".join(args.config_file_list)
|
||||
|
||||
main.call(config_file_list_str=config_file_list_str, recover=args.recover, name=args.name)
|
||||
91
scripts/convert_lora_to_peft_format.py
Normal file
91
scripts/convert_lora_to_peft_format.py
Normal file
@@ -0,0 +1,91 @@
|
||||
# currently only works with flux as support is not quite there yet
|
||||
|
||||
import argparse
|
||||
import os.path
|
||||
from collections import OrderedDict
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
'input_path',
|
||||
type=str,
|
||||
help='Path to original sdxl model'
|
||||
)
|
||||
parser.add_argument(
|
||||
'output_path',
|
||||
type=str,
|
||||
help='output path'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
args.input_path = os.path.abspath(args.input_path)
|
||||
args.output_path = os.path.abspath(args.output_path)
|
||||
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
|
||||
state_dict = load_file(args.input_path)
|
||||
|
||||
# peft doesnt have an alpha so we need to scale the weights
|
||||
alpha_keys = [
|
||||
'lora_transformer_single_transformer_blocks_0_attn_to_q.alpha' # flux
|
||||
]
|
||||
|
||||
# keys where the rank is in the first dimension
|
||||
rank_idx0_keys = [
|
||||
'lora_transformer_single_transformer_blocks_0_attn_to_q.lora_down.weight'
|
||||
# 'transformer.single_transformer_blocks.0.attn.to_q.lora_A.weight'
|
||||
]
|
||||
|
||||
alpha = None
|
||||
rank = None
|
||||
|
||||
for key in rank_idx0_keys:
|
||||
if key in state_dict:
|
||||
rank = int(state_dict[key].shape[0])
|
||||
break
|
||||
|
||||
if rank is None:
|
||||
raise ValueError(f'Could not find rank in state dict')
|
||||
|
||||
for key in alpha_keys:
|
||||
if key in state_dict:
|
||||
alpha = int(state_dict[key])
|
||||
break
|
||||
|
||||
if alpha is None:
|
||||
# set to rank if not found
|
||||
alpha = rank
|
||||
|
||||
|
||||
up_multiplier = alpha / rank
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
if key.endswith('.alpha'):
|
||||
continue
|
||||
|
||||
orig_dtype = value.dtype
|
||||
|
||||
new_val = value.float() * up_multiplier
|
||||
|
||||
new_key = key
|
||||
new_key = new_key.replace('lora_transformer_', 'transformer.')
|
||||
for i in range(100):
|
||||
new_key = new_key.replace(f'transformer_blocks_{i}_', f'transformer_blocks.{i}.')
|
||||
new_key = new_key.replace('lora_down', 'lora_A')
|
||||
new_key = new_key.replace('lora_up', 'lora_B')
|
||||
new_key = new_key.replace('_lora', '.lora')
|
||||
new_key = new_key.replace('attn_', 'attn.')
|
||||
new_key = new_key.replace('ff_', 'ff.')
|
||||
new_key = new_key.replace('context_net_', 'context.net.')
|
||||
new_key = new_key.replace('0_proj', '0.proj')
|
||||
new_key = new_key.replace('norm_linear', 'norm.linear')
|
||||
new_key = new_key.replace('norm_out_linear', 'norm_out.linear')
|
||||
new_key = new_key.replace('to_out_', 'to_out.')
|
||||
|
||||
new_state_dict[new_key] = new_val.to(orig_dtype)
|
||||
|
||||
save_file(new_state_dict, args.output_path, meta)
|
||||
print(f'Saved to {args.output_path}')
|
||||
20
scripts/generate_sampler_step_scales.py
Normal file
20
scripts/generate_sampler_step_scales.py
Normal file
@@ -0,0 +1,20 @@
|
||||
import argparse
|
||||
import torch
|
||||
import os
|
||||
from diffusers import StableDiffusionPipeline
|
||||
import sys
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
# add project root to path
|
||||
sys.path.append(PROJECT_ROOT)
|
||||
|
||||
SAMPLER_SCALES_ROOT = os.path.join(PROJECT_ROOT, 'toolkit', 'samplers_scales')
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description='Process some images.')
|
||||
add_arg = parser.add_argument
|
||||
add_arg('--model', type=str, required=True, help='Path to model')
|
||||
add_arg('--sampler', type=str, required=True, help='Name of sampler')
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
import sys
|
||||
import os
|
||||
ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
sys.path.append(ROOT_DIR)
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
42
scripts/patch_te_adapter.py
Normal file
42
scripts/patch_te_adapter.py
Normal file
@@ -0,0 +1,42 @@
|
||||
import torch
|
||||
from safetensors.torch import save_file, load_file
|
||||
from collections import OrderedDict
|
||||
meta = OrderedDict()
|
||||
meta["format"] ="pt"
|
||||
|
||||
attn_dict = load_file("/mnt/Train/out/ip_adapter/sd15_bigG/sd15_bigG_000266000.safetensors")
|
||||
state_dict = load_file("/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors")
|
||||
|
||||
attn_list = []
|
||||
for key, value in state_dict.items():
|
||||
if "attn1" in key:
|
||||
attn_list.append(key)
|
||||
|
||||
attn_names = ['down_blocks.0.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.0.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.2.transformer_blocks.0.attn2.processor', 'mid_block.attentions.0.transformer_blocks.0.attn2.processor']
|
||||
|
||||
adapter_names = []
|
||||
for i in range(100):
|
||||
if f'te_adapter.adapter_modules.{i}.to_k_adapter.weight' in attn_dict:
|
||||
adapter_names.append(f"te_adapter.adapter_modules.{i}.adapter")
|
||||
|
||||
|
||||
for i in range(len(adapter_names)):
|
||||
adapter_name = adapter_names[i]
|
||||
attn_name = attn_names[i]
|
||||
adapter_k_name = adapter_name[:-8] + '.to_k_adapter.weight'
|
||||
adapter_v_name = adapter_name[:-8] + '.to_v_adapter.weight'
|
||||
state_k_name = attn_name.replace(".processor", ".to_k.weight")
|
||||
state_v_name = attn_name.replace(".processor", ".to_v.weight")
|
||||
if adapter_k_name in attn_dict:
|
||||
state_dict[state_k_name] = attn_dict[adapter_k_name]
|
||||
state_dict[state_v_name] = attn_dict[adapter_v_name]
|
||||
else:
|
||||
print("adapter_k_name", adapter_k_name)
|
||||
print("state_k_name", state_k_name)
|
||||
|
||||
for key, value in state_dict.items():
|
||||
state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
save_file(state_dict, "/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors", metadata=meta)
|
||||
|
||||
print("Done")
|
||||
65
scripts/repair_dataset_folder.py
Normal file
65
scripts/repair_dataset_folder.py
Normal file
@@ -0,0 +1,65 @@
|
||||
import argparse
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from tqdm import tqdm
|
||||
import os
|
||||
|
||||
parser = argparse.ArgumentParser(description='Process some images.')
|
||||
parser.add_argument("input_folder", type=str, help="Path to folder containing images")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
img_types = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
|
||||
# find all images in the input folder
|
||||
images = []
|
||||
for root, _, files in os.walk(args.input_folder):
|
||||
for file in files:
|
||||
if file.lower().endswith(tuple(img_types)):
|
||||
images.append(os.path.join(root, file))
|
||||
print(f"Found {len(images)} images")
|
||||
|
||||
num_skipped = 0
|
||||
num_repaired = 0
|
||||
num_deleted = 0
|
||||
|
||||
pbar = tqdm(total=len(images), desc=f"Repaired {num_repaired} images", unit="image")
|
||||
for img_path in images:
|
||||
filename = os.path.basename(img_path)
|
||||
filename_no_ext, file_extension = os.path.splitext(filename)
|
||||
# if it is jpg, ignore
|
||||
if file_extension.lower() == '.jpg':
|
||||
num_skipped += 1
|
||||
pbar.update(1)
|
||||
|
||||
continue
|
||||
|
||||
try:
|
||||
img = Image.open(img_path)
|
||||
except Exception as e:
|
||||
print(f"Error opening {img_path}: {e}")
|
||||
# delete it
|
||||
os.remove(img_path)
|
||||
num_deleted += 1
|
||||
pbar.update(1)
|
||||
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
|
||||
continue
|
||||
|
||||
|
||||
try:
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error rotating {img_path}: {e}")
|
||||
|
||||
new_path = os.path.join(os.path.dirname(img_path), filename_no_ext + '.jpg')
|
||||
|
||||
img = img.convert("RGB")
|
||||
img.save(new_path, quality=95)
|
||||
# remove the old file
|
||||
os.remove(img_path)
|
||||
num_repaired += 1
|
||||
pbar.update(1)
|
||||
# update pbar
|
||||
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
|
||||
|
||||
print("Done")
|
||||
@@ -54,6 +54,7 @@ parser.add_argument('--name', type=str, default='stable_diffusion', help='name f
|
||||
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
|
||||
parser.add_argument('--refiner', action='store_true', help='is refiner model')
|
||||
parser.add_argument('--ssd', action='store_true', help='is ssd model')
|
||||
parser.add_argument('--vega', action='store_true', help='is vega model')
|
||||
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
|
||||
|
||||
args = parser.parse_args()
|
||||
@@ -66,15 +67,15 @@ print(f'Loading diffusers model')
|
||||
|
||||
ignore_ldm_begins_with = []
|
||||
|
||||
diffusers_file_path = file_path
|
||||
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
|
||||
if args.ssd:
|
||||
diffusers_file_path = "segmind/SSD-1B"
|
||||
if args.vega:
|
||||
diffusers_file_path = "segmind/Segmind-Vega"
|
||||
|
||||
# if args.refiner:
|
||||
# diffusers_file_path = "stabilityai/stable-diffusion-xl-refiner-1.0"
|
||||
|
||||
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
|
||||
|
||||
if not args.refiner:
|
||||
|
||||
diffusers_model_config = ModelConfig(
|
||||
@@ -82,6 +83,7 @@ if not args.refiner:
|
||||
is_xl=args.sdxl,
|
||||
is_v2=args.sd2,
|
||||
is_ssd=args.ssd,
|
||||
is_vega=args.vega,
|
||||
dtype=dtype,
|
||||
)
|
||||
diffusers_sd = StableDiffusion(
|
||||
@@ -157,7 +159,7 @@ te_suffix = ''
|
||||
proj_pattern_weight = None
|
||||
proj_pattern_bias = None
|
||||
text_proj_layer = None
|
||||
if args.sdxl or args.ssd:
|
||||
if args.sdxl or args.ssd or args.vega:
|
||||
te_suffix = '1'
|
||||
ldm_res_block_prefix = "conditioner.embedders.1.model.transformer.resblocks"
|
||||
proj_pattern_weight = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
|
||||
@@ -176,10 +178,13 @@ if args.sd2:
|
||||
proj_pattern_bias = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
|
||||
text_proj_layer = "cond_stage_model.model.text_projection"
|
||||
|
||||
if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
if args.sdxl or args.sd2 or args.ssd or args.refiner or args.vega:
|
||||
if "conditioner.embedders.1.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection"].shape[0])
|
||||
elif "conditioner.embedders.1.model.text_projection.weight" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection.weight"].shape[0])
|
||||
elif "conditioner.embedders.0.model.text_projection" in ldm_dict_keys:
|
||||
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
|
||||
d_model = int(ldm_state_dict["conditioner.embedders.0.model.text_projection"].shape[0])
|
||||
@@ -191,6 +196,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
try:
|
||||
match = re.match(proj_pattern_weight, ldm_key)
|
||||
if match:
|
||||
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight":
|
||||
print("here")
|
||||
number = int(match.group(1))
|
||||
new_val = torch.cat([
|
||||
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"],
|
||||
@@ -217,6 +224,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
],
|
||||
}
|
||||
|
||||
matched_ldm_keys.append(ldm_key)
|
||||
|
||||
# text_model_dict[new_key + ".q_proj.weight"] = checkpoint[key][:d_model, :]
|
||||
# text_model_dict[new_key + ".k_proj.weight"] = checkpoint[key][d_model: d_model * 2, :]
|
||||
# text_model_dict[new_key + ".v_proj.weight"] = checkpoint[key][d_model * 2:, :]
|
||||
@@ -266,6 +275,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
|
||||
],
|
||||
}
|
||||
|
||||
matched_ldm_keys.append(ldm_key)
|
||||
|
||||
# add diffusers operators
|
||||
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias"] = {
|
||||
"slice": [
|
||||
@@ -298,6 +309,9 @@ for ldm_key in ldm_dict_keys:
|
||||
ldm_shape_tuple = ldm_state_dict[ldm_key].shape
|
||||
ldm_reduced_shape_tuple = get_reduced_shape(ldm_shape_tuple)
|
||||
for diffusers_key in diffusers_dict_keys:
|
||||
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight" and diffusers_key == "te1_text_model.encoder.layers.0.self_attn.q_proj.weight":
|
||||
print("here")
|
||||
|
||||
diffusers_shape_tuple = diffusers_state_dict[diffusers_key].shape
|
||||
diffusers_reduced_shape_tuple = get_reduced_shape(diffusers_shape_tuple)
|
||||
|
||||
@@ -356,6 +370,8 @@ if args.sdxl:
|
||||
name += '_sdxl'
|
||||
elif args.ssd:
|
||||
name += '_ssd'
|
||||
elif args.vega:
|
||||
name += '_vega'
|
||||
elif args.refiner:
|
||||
name += '_refiner'
|
||||
elif args.sd2:
|
||||
|
||||
180
testing/merge_in_text_encoder_adapter.py
Normal file
180
testing/merge_in_text_encoder_adapter.py
Normal file
@@ -0,0 +1,180 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
from diffusers import StableDiffusionPipeline, UNet2DConditionModel, PixArtSigmaPipeline, Transformer2DModel, PixArtTransformer2DModel
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
|
||||
# model_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_v01_000527000"
|
||||
# te_path = "google/flan-t5-xl"
|
||||
# te_aug_path = "/mnt/Train/out/ip_adapter/t5xx_sd15_v1/t5xx_sd15_v1_000032000.safetensors"
|
||||
# output_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_t5xl_raw"
|
||||
model_path = "/home/jaret/Dev/models/hf/objective-reality-16ch"
|
||||
te_path = "google/flan-t5-xl"
|
||||
te_aug_path = "/mnt/Train2/out/ip_adapter/t5xl-sd15-16ch_v1/t5xl-sd15-16ch_v1_000115000.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/t5xl-sd15-16ch_sd15_v1"
|
||||
|
||||
|
||||
print("Loading te adapter")
|
||||
te_aug_sd = load_file(te_aug_path)
|
||||
|
||||
print("Loading model")
|
||||
is_diffusers = (not os.path.exists(model_path)) or os.path.isdir(model_path)
|
||||
|
||||
# if "pixart" in model_path.lower():
|
||||
is_pixart = "pixart" in model_path.lower()
|
||||
|
||||
pipeline_class = StableDiffusionPipeline
|
||||
|
||||
# transformer = PixArtTransformer2DModel.from_pretrained('PixArt-alpha/PixArt-Sigma-XL-2-512-MS', subfolder='transformer', torch_dtype=torch.float16)
|
||||
|
||||
if is_pixart:
|
||||
pipeline_class = PixArtSigmaPipeline
|
||||
|
||||
if is_diffusers:
|
||||
sd = pipeline_class.from_pretrained(model_path, torch_dtype=torch.float16)
|
||||
else:
|
||||
sd = pipeline_class.from_single_file(model_path, torch_dtype=torch.float16)
|
||||
|
||||
print("Loading Text Encoder")
|
||||
# Load the text encoder
|
||||
te = T5EncoderModel.from_pretrained(te_path, torch_dtype=torch.float16)
|
||||
|
||||
# patch it
|
||||
sd.text_encoder = te
|
||||
sd.tokenizer = T5Tokenizer.from_pretrained(te_path)
|
||||
|
||||
if is_pixart:
|
||||
unet = sd.transformer
|
||||
unet_sd = sd.transformer.state_dict()
|
||||
else:
|
||||
unet = sd.unet
|
||||
unet_sd = sd.unet.state_dict()
|
||||
|
||||
|
||||
if is_pixart:
|
||||
weight_idx = 0
|
||||
else:
|
||||
weight_idx = 1
|
||||
|
||||
new_cross_attn_dim = None
|
||||
|
||||
# count the num of params in state dict
|
||||
start_params = sum([v.numel() for v in unet_sd.values()])
|
||||
|
||||
print("Building")
|
||||
attn_processor_keys = []
|
||||
if is_pixart:
|
||||
transformer: Transformer2DModel = unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
|
||||
# cross attention
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
|
||||
else:
|
||||
attn_processor_keys = list(unet.attn_processors.keys())
|
||||
|
||||
for name in attn_processor_keys:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") or name.endswith(
|
||||
"attn1") else \
|
||||
unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
hidden_size = unet.config['cross_attention_dim']
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
pass
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
|
||||
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
|
||||
|
||||
te_aug_name = None
|
||||
while True:
|
||||
if is_pixart:
|
||||
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
|
||||
else:
|
||||
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
|
||||
if f"{te_aug_name}.weight" in te_aug_sd:
|
||||
# increment so we dont redo it next time
|
||||
weight_idx += 1
|
||||
break
|
||||
else:
|
||||
weight_idx += 1
|
||||
|
||||
if weight_idx > 1000:
|
||||
raise ValueError("Could not find the next weight")
|
||||
|
||||
orig_weight_shape_k = list(unet_sd[layer_name + ".to_k.weight"].shape)
|
||||
new_weight_shape_k = list(te_aug_sd[te_aug_name + ".weight"].shape)
|
||||
orig_weight_shape_v = list(unet_sd[layer_name + ".to_v.weight"].shape)
|
||||
new_weight_shape_v = list(te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"].shape)
|
||||
|
||||
unet_sd[layer_name + ".to_k.weight"] = te_aug_sd[te_aug_name + ".weight"]
|
||||
unet_sd[layer_name + ".to_v.weight"] = te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"]
|
||||
|
||||
if new_cross_attn_dim is None:
|
||||
new_cross_attn_dim = unet_sd[layer_name + ".to_k.weight"].shape[1]
|
||||
|
||||
|
||||
|
||||
if is_pixart:
|
||||
# copy the caption_projection weight
|
||||
del unet_sd['caption_projection.linear_1.bias']
|
||||
del unet_sd['caption_projection.linear_1.weight']
|
||||
del unet_sd['caption_projection.linear_2.bias']
|
||||
del unet_sd['caption_projection.linear_2.weight']
|
||||
|
||||
print("Saving unmodified model")
|
||||
sd = sd.to("cpu", torch.float16)
|
||||
sd.save_pretrained(
|
||||
output_path,
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
# overwrite the unet
|
||||
if is_pixart:
|
||||
unet_folder = os.path.join(output_path, "transformer")
|
||||
else:
|
||||
unet_folder = os.path.join(output_path, "unet")
|
||||
|
||||
# move state_dict to cpu
|
||||
unet_sd = {k: v.clone().cpu().to(torch.float16) for k, v in unet_sd.items()}
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
print("Patching")
|
||||
|
||||
save_file(unet_sd, os.path.join(unet_folder, "diffusion_pytorch_model.safetensors"), meta)
|
||||
|
||||
# load the json file
|
||||
with open(os.path.join(unet_folder, "config.json"), 'r') as f:
|
||||
config = json.load(f)
|
||||
|
||||
config['cross_attention_dim'] = new_cross_attn_dim
|
||||
|
||||
if is_pixart:
|
||||
config['caption_channels'] = None
|
||||
|
||||
# save it
|
||||
with open(os.path.join(unet_folder, "config.json"), 'w') as f:
|
||||
json.dump(config, f, indent=2)
|
||||
|
||||
print("Done")
|
||||
|
||||
new_params = sum([v.numel() for v in unet_sd.values()])
|
||||
|
||||
# print new and old params with , formatted
|
||||
print(f"Old params: {start_params:,}")
|
||||
print(f"New params: {new_params:,}")
|
||||
62
testing/shrink_pixart.py
Normal file
62
testing/shrink_pixart.py
Normal file
@@ -0,0 +1,62 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
|
||||
|
||||
state_dict = load_file(model_path)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
# Move non-blocks over
|
||||
for key, value in state_dict.items():
|
||||
if not key.startswith("transformer_blocks."):
|
||||
new_state_dict[key] = value
|
||||
|
||||
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
|
||||
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
|
||||
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
|
||||
'transformer_blocks.{idx}.scale_shift_table']
|
||||
|
||||
# New block idx 0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27
|
||||
|
||||
current_idx = 0
|
||||
for i in range(28):
|
||||
if i not in [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]:
|
||||
# todo merge in with previous block
|
||||
for name in block_names:
|
||||
try:
|
||||
new_state_dict_key = name.format(idx=current_idx - 1)
|
||||
old_state_dict_key = name.format(idx=i)
|
||||
new_state_dict[new_state_dict_key] = (new_state_dict[new_state_dict_key] * 0.5) + (state_dict[old_state_dict_key] * 0.5)
|
||||
except KeyError:
|
||||
raise KeyError(f"KeyError: {name.format(idx=current_idx)}")
|
||||
else:
|
||||
for name in block_names:
|
||||
new_state_dict[name.format(idx=current_idx)] = state_dict[name.format(idx=i)]
|
||||
current_idx += 1
|
||||
|
||||
|
||||
# make sure they are all fp16 and on cpu
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.to(torch.float16).cpu()
|
||||
|
||||
# save the new state dict
|
||||
save_file(new_state_dict, output_path, metadata=meta)
|
||||
|
||||
new_param_count = sum([v.numel() for v in new_state_dict.values()])
|
||||
old_param_count = sum([v.numel() for v in state_dict.values()])
|
||||
|
||||
print(f"Old param count: {old_param_count:,}")
|
||||
print(f"New param count: {new_param_count:,}")
|
||||
81
testing/shrink_pixart2.py
Normal file
81
testing/shrink_pixart2.py
Normal file
@@ -0,0 +1,81 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
|
||||
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
|
||||
|
||||
state_dict = load_file(model_path)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta["format"] = "pt"
|
||||
|
||||
new_state_dict = {}
|
||||
|
||||
# Move non-blocks over
|
||||
for key, value in state_dict.items():
|
||||
if not key.startswith("transformer_blocks."):
|
||||
new_state_dict[key] = value
|
||||
|
||||
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
|
||||
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
|
||||
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
|
||||
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
|
||||
'transformer_blocks.{idx}.scale_shift_table']
|
||||
|
||||
# Blocks to keep
|
||||
# keep_blocks = [0, 1, 2, 6, 10, 14, 18, 22, 26, 27]
|
||||
keep_blocks = [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]
|
||||
|
||||
|
||||
def weighted_merge(kept_block, removed_block, weight):
|
||||
return kept_block * (1 - weight) + removed_block * weight
|
||||
|
||||
|
||||
# First, copy all kept blocks to new_state_dict
|
||||
for i, old_idx in enumerate(keep_blocks):
|
||||
for name in block_names:
|
||||
old_key = name.format(idx=old_idx)
|
||||
new_key = name.format(idx=i)
|
||||
new_state_dict[new_key] = state_dict[old_key].clone()
|
||||
|
||||
# Then, merge information from removed blocks
|
||||
for i in range(28):
|
||||
if i not in keep_blocks:
|
||||
# Find the nearest kept blocks
|
||||
prev_kept = max([b for b in keep_blocks if b < i])
|
||||
next_kept = min([b for b in keep_blocks if b > i])
|
||||
|
||||
# Calculate the weight based on position
|
||||
weight = (i - prev_kept) / (next_kept - prev_kept)
|
||||
|
||||
for name in block_names:
|
||||
removed_key = name.format(idx=i)
|
||||
prev_new_key = name.format(idx=keep_blocks.index(prev_kept))
|
||||
next_new_key = name.format(idx=keep_blocks.index(next_kept))
|
||||
|
||||
# Weighted merge for previous kept block
|
||||
new_state_dict[prev_new_key] = weighted_merge(new_state_dict[prev_new_key], state_dict[removed_key], weight)
|
||||
|
||||
# Weighted merge for next kept block
|
||||
new_state_dict[next_new_key] = weighted_merge(new_state_dict[next_new_key], state_dict[removed_key],
|
||||
1 - weight)
|
||||
|
||||
# Convert to fp16 and move to CPU
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.to(torch.float16).cpu()
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict, output_path, metadata=meta)
|
||||
|
||||
new_param_count = sum([v.numel() for v in new_state_dict.values()])
|
||||
old_param_count = sum([v.numel() for v in state_dict.values()])
|
||||
|
||||
print(f"Old param count: {old_param_count:,}")
|
||||
print(f"New param count: {new_param_count:,}")
|
||||
84
testing/shrink_pixart_sm.py
Normal file
84
testing/shrink_pixart_sm.py
Normal file
@@ -0,0 +1,84 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
original_shape = weight.shape
|
||||
flattened = weight.view(-1, original_shape[-1])
|
||||
|
||||
if flattened.shape[1] <= target_size:
|
||||
return weight
|
||||
|
||||
U, S, V = torch.svd(flattened)
|
||||
reduced = torch.mm(U[:, :target_size], torch.diag(S[:target_size]))
|
||||
|
||||
if reduced.shape[1] < target_size:
|
||||
padding = torch.zeros(reduced.shape[0], target_size - reduced.shape[1], device=device)
|
||||
reduced = torch.cat((reduced, padding), dim=1)
|
||||
|
||||
return reduced.view(original_shape[:-1] + (target_size,))
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.to(device, torch.float32)
|
||||
original_size = bias.shape[0]
|
||||
|
||||
if original_size <= target_size:
|
||||
return torch.nn.functional.pad(bias, (0, target_size - original_size))
|
||||
else:
|
||||
return bias.view(-1, original_size // target_size).mean(dim=1)[:target_size]
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
source_hidden_size = 1152
|
||||
target_hidden_size = 1024
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == source_hidden_size:
|
||||
value = value[:target_hidden_size]
|
||||
elif value.shape[0] == source_hidden_size * 4:
|
||||
value = value[:target_hidden_size * 4]
|
||||
elif value.shape[0] == source_hidden_size * 6:
|
||||
value = value[:target_hidden_size * 6]
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == source_hidden_size and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = value[:, :target_hidden_size]
|
||||
elif len(value.shape) > 1 and value.shape[1] == source_hidden_size * 4:
|
||||
value = value[:, :target_hidden_size * 4]
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == source_hidden_size:
|
||||
value = value[:target_hidden_size]
|
||||
elif value.shape[0] == source_hidden_size * 4:
|
||||
value = value[:target_hidden_size * 4]
|
||||
elif value.shape[0] == source_hidden_size * 6:
|
||||
value = value[:target_hidden_size * 6]
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
110
testing/shrink_pixart_sm2.py
Normal file
110
testing/shrink_pixart_sm2.py
Normal file
@@ -0,0 +1,110 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
original_shape = weight.shape
|
||||
|
||||
if len(original_shape) == 1:
|
||||
# For 1D tensors, simply truncate
|
||||
return weight[:target_size]
|
||||
|
||||
if original_shape[0] <= target_size:
|
||||
return weight
|
||||
|
||||
# Reshape the tensor to 2D
|
||||
flattened = weight.reshape(original_shape[0], -1)
|
||||
|
||||
# Perform SVD
|
||||
U, S, V = torch.svd(flattened)
|
||||
|
||||
# Reduce the dimensions
|
||||
reduced = torch.mm(U[:target_size, :], torch.diag(S)).mm(V.t())
|
||||
|
||||
# Reshape back to the original shape with reduced first dimension
|
||||
new_shape = (target_size,) + original_shape[1:]
|
||||
return reduced.reshape(new_shape)
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.to(device, torch.float32)
|
||||
return bias[:target_size]
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == 1152:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
|
||||
# reshape to (1152, -1)
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 512)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
# value = reduce_weight(value.t(), 576).t().contiguous()
|
||||
value = reduce_weight(value, 512)
|
||||
pass
|
||||
elif value.shape[0] == 4608:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 2048)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 3072)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 3072)
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
elif len(value.shape) > 1 and value.shape[1] == 4608:
|
||||
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == 1152:
|
||||
value = reduce_bias(value, 512)
|
||||
elif value.shape[0] == 4608:
|
||||
value = reduce_bias(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
value = reduce_bias(value, 3072)
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
100
testing/shrink_pixart_sm3.py
Normal file
100
testing/shrink_pixart_sm3.py
Normal file
@@ -0,0 +1,100 @@
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from collections import OrderedDict
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = "pt"
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
def reduce_weight(weight, target_size):
|
||||
weight = weight.to(device, torch.float32)
|
||||
# resize so target_size is the first dimension
|
||||
tmp_weight = weight.view(1, 1, weight.shape[0], weight.shape[1])
|
||||
|
||||
# use interpolate to resize the tensor
|
||||
new_weight = torch.nn.functional.interpolate(tmp_weight, size=(target_size, weight.shape[1]), mode='bicubic', align_corners=True)
|
||||
|
||||
# reshape back to original shape
|
||||
return new_weight.view(target_size, weight.shape[1])
|
||||
|
||||
|
||||
def reduce_bias(bias, target_size):
|
||||
bias = bias.view(1, 1, bias.shape[0], 1)
|
||||
|
||||
new_bias = torch.nn.functional.interpolate(bias, size=(target_size, 1), mode='bicubic', align_corners=True)
|
||||
|
||||
return new_bias.view(target_size)
|
||||
|
||||
|
||||
# Load your original state dict
|
||||
state_dict = load_file(
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
|
||||
|
||||
# Create a new state dict for the reduced model
|
||||
new_state_dict = {}
|
||||
|
||||
for key, value in state_dict.items():
|
||||
value = value.to(device, torch.float32)
|
||||
|
||||
if 'weight' in key or 'scale_shift_table' in key:
|
||||
if value.shape[0] == 1152:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
|
||||
# reshape to (1152, -1)
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 512)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
# value = reduce_weight(value.t(), 576).t().contiguous()
|
||||
value = reduce_weight(value, 512)
|
||||
pass
|
||||
elif value.shape[0] == 4608:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 2048)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
if len(value.shape) == 4:
|
||||
orig_shape = value.shape
|
||||
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
|
||||
value = value.view(value.shape[0], -1)
|
||||
value = reduce_weight(value, 3072)
|
||||
value = value.view(output_shape)
|
||||
else:
|
||||
value = reduce_weight(value, 3072)
|
||||
|
||||
if len(value.shape) > 1 and value.shape[
|
||||
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
|
||||
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
elif len(value.shape) > 1 and value.shape[1] == 4608:
|
||||
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
|
||||
pass
|
||||
|
||||
elif 'bias' in key:
|
||||
if value.shape[0] == 1152:
|
||||
value = reduce_bias(value, 512)
|
||||
elif value.shape[0] == 4608:
|
||||
value = reduce_bias(value, 2048)
|
||||
elif value.shape[0] == 6912:
|
||||
value = reduce_bias(value, 3072)
|
||||
|
||||
new_state_dict[key] = value
|
||||
|
||||
# Move all to CPU and convert to float16
|
||||
for key, value in new_state_dict.items():
|
||||
new_state_dict[key] = value.cpu().to(torch.float16)
|
||||
|
||||
# Save the new state dict
|
||||
save_file(new_state_dict,
|
||||
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
|
||||
metadata=meta)
|
||||
|
||||
print("Done!")
|
||||
@@ -7,11 +7,13 @@ from torchvision import transforms
|
||||
import sys
|
||||
import os
|
||||
import cv2
|
||||
import random
|
||||
from transformers import CLIPImageProcessor
|
||||
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
|
||||
from toolkit.image_utils import show_img
|
||||
import torchvision.transforms.functional
|
||||
from toolkit.image_utils import show_img, show_tensors
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
@@ -21,12 +23,14 @@ from toolkit.data_loader import AiToolkitDataset, get_dataloader_from_datasets,
|
||||
trigger_dataloader_setup_epoch
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
import argparse
|
||||
from tqdm import tqdm
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('dataset_folder', type=str, default='input')
|
||||
parser.add_argument('--epochs', type=int, default=1)
|
||||
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
dataset_folder = args.dataset_folder
|
||||
@@ -34,70 +38,87 @@ resolution = 1024
|
||||
bucket_tolerance = 64
|
||||
batch_size = 1
|
||||
|
||||
clip_processor = CLIPImageProcessor.from_pretrained("openai/clip-vit-base-patch16")
|
||||
|
||||
class FakeAdapter:
|
||||
def __init__(self):
|
||||
self.clip_image_processor = clip_processor
|
||||
|
||||
|
||||
## make fake sd
|
||||
class FakeSD:
|
||||
def __init__(self):
|
||||
self.adapter = FakeAdapter()
|
||||
|
||||
|
||||
|
||||
##
|
||||
|
||||
dataset_config = DatasetConfig(
|
||||
dataset_path=dataset_folder,
|
||||
# clip_image_path=dataset_folder,
|
||||
# square_crop=True,
|
||||
resolution=resolution,
|
||||
caption_ext='json',
|
||||
# caption_ext='json',
|
||||
default_caption='default',
|
||||
# clip_image_path='/mnt/Datasets2/regs/yetibear_xl_v14/random_aspect/',
|
||||
buckets=True,
|
||||
bucket_tolerance=bucket_tolerance,
|
||||
poi='person',
|
||||
augmentations=[
|
||||
{
|
||||
'method': 'RandomBrightnessContrast',
|
||||
'brightness_limit': (-0.3, 0.3),
|
||||
'contrast_limit': (-0.3, 0.3),
|
||||
'brightness_by_max': False,
|
||||
'p': 1.0
|
||||
},
|
||||
{
|
||||
'method': 'HueSaturationValue',
|
||||
'hue_shift_limit': (-0, 0),
|
||||
'sat_shift_limit': (-40, 40),
|
||||
'val_shift_limit': (-40, 40),
|
||||
'p': 1.0
|
||||
},
|
||||
# {
|
||||
# 'method': 'RGBShift',
|
||||
# 'r_shift_limit': (-20, 20),
|
||||
# 'g_shift_limit': (-20, 20),
|
||||
# 'b_shift_limit': (-20, 20),
|
||||
# 'p': 1.0
|
||||
# },
|
||||
]
|
||||
|
||||
|
||||
# poi='person',
|
||||
# shuffle_augmentations=True,
|
||||
# augmentations=[
|
||||
# {
|
||||
# 'method': 'Posterize',
|
||||
# 'num_bits': [(0, 4), (0, 4), (0, 4)],
|
||||
# 'p': 1.0
|
||||
# },
|
||||
#
|
||||
# ]
|
||||
)
|
||||
|
||||
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size)
|
||||
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size, sd=FakeSD())
|
||||
|
||||
|
||||
# run through an epoch ang check sizes
|
||||
dataloader_iterator = iter(dataloader)
|
||||
for epoch in range(args.epochs):
|
||||
for batch in dataloader:
|
||||
for batch in tqdm(dataloader):
|
||||
batch: 'DataLoaderBatchDTO'
|
||||
img_batch = batch.tensor
|
||||
batch_size, channels, height, width = img_batch.shape
|
||||
|
||||
chunks = torch.chunk(img_batch, batch_size, dim=0)
|
||||
# put them so they are size by side
|
||||
big_img = torch.cat(chunks, dim=3)
|
||||
big_img = big_img.squeeze(0)
|
||||
# img_batch = color_block_imgs(img_batch, neg1_1=True)
|
||||
|
||||
min_val = big_img.min()
|
||||
max_val = big_img.max()
|
||||
# chunks = torch.chunk(img_batch, batch_size, dim=0)
|
||||
# # put them so they are size by side
|
||||
# big_img = torch.cat(chunks, dim=3)
|
||||
# big_img = big_img.squeeze(0)
|
||||
#
|
||||
# control_chunks = torch.chunk(batch.clip_image_tensor, batch_size, dim=0)
|
||||
# big_control_img = torch.cat(control_chunks, dim=3)
|
||||
# big_control_img = big_control_img.squeeze(0) * 2 - 1
|
||||
#
|
||||
#
|
||||
# # resize control image
|
||||
# big_control_img = torchvision.transforms.Resize((width, height))(big_control_img)
|
||||
#
|
||||
# big_img = torch.cat([big_img, big_control_img], dim=2)
|
||||
#
|
||||
# min_val = big_img.min()
|
||||
# max_val = big_img.max()
|
||||
#
|
||||
# big_img = (big_img / 2 + 0.5).clamp(0, 1)
|
||||
|
||||
big_img = (big_img / 2 + 0.5).clamp(0, 1)
|
||||
big_img = img_batch
|
||||
# big_img = big_img.clamp(-1, 1)
|
||||
|
||||
show_tensors(big_img)
|
||||
|
||||
# convert to image
|
||||
img = transforms.ToPILImage()(big_img)
|
||||
# img = transforms.ToPILImage()(big_img)
|
||||
#
|
||||
# show_img(img)
|
||||
|
||||
show_img(img)
|
||||
|
||||
time.sleep(1.0)
|
||||
time.sleep(0.2)
|
||||
# if not last epoch
|
||||
if epoch < args.epochs - 1:
|
||||
trigger_dataloader_setup_epoch(dataloader)
|
||||
|
||||
113
testing/test_vae.py
Normal file
113
testing/test_vae.py
Normal file
@@ -0,0 +1,113 @@
|
||||
import argparse
|
||||
import os
|
||||
from PIL import Image
|
||||
import torch
|
||||
from torchvision.transforms import Resize, ToTensor
|
||||
from diffusers import AutoencoderKL
|
||||
from pytorch_fid import fid_score
|
||||
from skimage.metrics import peak_signal_noise_ratio as psnr
|
||||
import lpips
|
||||
from tqdm import tqdm
|
||||
from torchvision import transforms
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
def load_images(folder_path):
|
||||
images = []
|
||||
for filename in os.listdir(folder_path):
|
||||
if filename.lower().endswith(('.png', '.jpg', '.jpeg')):
|
||||
img_path = os.path.join(folder_path, filename)
|
||||
images.append(img_path)
|
||||
return images
|
||||
|
||||
|
||||
def paramiter_count(model):
|
||||
state_dict = model.state_dict()
|
||||
paramiter_count = 0
|
||||
for key in state_dict:
|
||||
paramiter_count += torch.numel(state_dict[key])
|
||||
return int(paramiter_count)
|
||||
|
||||
|
||||
def calculate_metrics(vae, images, max_imgs=-1):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
vae = vae.to(device)
|
||||
lpips_model = lpips.LPIPS(net='alex').to(device)
|
||||
|
||||
rfid_scores = []
|
||||
psnr_scores = []
|
||||
lpips_scores = []
|
||||
|
||||
# transform = transforms.Compose([
|
||||
# transforms.Resize(256, antialias=True),
|
||||
# transforms.CenterCrop(256)
|
||||
# ])
|
||||
# needs values between -1 and 1
|
||||
to_tensor = ToTensor()
|
||||
|
||||
if max_imgs > 0 and len(images) > max_imgs:
|
||||
images = images[:max_imgs]
|
||||
|
||||
for img_path in tqdm(images):
|
||||
try:
|
||||
img = Image.open(img_path).convert('RGB')
|
||||
# img_tensor = to_tensor(transform(img)).unsqueeze(0).to(device)
|
||||
img_tensor = to_tensor(img).unsqueeze(0).to(device)
|
||||
img_tensor = 2 * img_tensor - 1
|
||||
# if width or height is not divisible by 8, crop it
|
||||
if img_tensor.shape[2] % 8 != 0 or img_tensor.shape[3] % 8 != 0:
|
||||
img_tensor = img_tensor[:, :, :img_tensor.shape[2] // 8 * 8, :img_tensor.shape[3] // 8 * 8]
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error processing {img_path}: {e}")
|
||||
continue
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
reconstructed = vae.decode(vae.encode(img_tensor).latent_dist.sample()).sample
|
||||
|
||||
# Calculate rFID
|
||||
# rfid = fid_score.calculate_frechet_distance(vae, img_tensor, reconstructed)
|
||||
# rfid_scores.append(rfid)
|
||||
|
||||
# Calculate PSNR
|
||||
psnr_val = psnr(img_tensor.cpu().numpy(), reconstructed.cpu().numpy())
|
||||
psnr_scores.append(psnr_val)
|
||||
|
||||
# Calculate LPIPS
|
||||
lpips_val = lpips_model(img_tensor, reconstructed).item()
|
||||
lpips_scores.append(lpips_val)
|
||||
|
||||
# avg_rfid = sum(rfid_scores) / len(rfid_scores)
|
||||
avg_rfid = 0
|
||||
avg_psnr = sum(psnr_scores) / len(psnr_scores)
|
||||
avg_lpips = sum(lpips_scores) / len(lpips_scores)
|
||||
|
||||
return avg_rfid, avg_psnr, avg_lpips
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Calculate average rFID, PSNR, and LPIPS for VAE reconstructions")
|
||||
parser.add_argument("--vae_path", type=str, required=True, help="Path to the VAE model")
|
||||
parser.add_argument("--image_folder", type=str, required=True, help="Path to the folder containing images")
|
||||
parser.add_argument("--max_imgs", type=int, default=-1, help="Max num of images. Default is -1 for all images.")
|
||||
args = parser.parse_args()
|
||||
|
||||
if os.path.isfile(args.vae_path):
|
||||
vae = AutoencoderKL.from_single_file(args.vae_path)
|
||||
else:
|
||||
vae = AutoencoderKL.from_pretrained(args.vae_path)
|
||||
vae.eval()
|
||||
vae = vae.to(device)
|
||||
print(f"Model has {paramiter_count(vae)} parameters")
|
||||
images = load_images(args.image_folder)
|
||||
|
||||
avg_rfid, avg_psnr, avg_lpips = calculate_metrics(vae, images, args.max_imgs)
|
||||
|
||||
# print(f"Average rFID: {avg_rfid}")
|
||||
print(f"Average PSNR: {avg_psnr}")
|
||||
print(f"Average LPIPS: {avg_lpips}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
55
toolkit/assistant_lora.py
Normal file
55
toolkit/assistant_lora.py
Normal file
@@ -0,0 +1,55 @@
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.config_modules import NetworkConfig
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from safetensors.torch import load_file
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
def load_assistant_lora_from_path(adapter_path, sd: 'StableDiffusion') -> LoRASpecialNetwork:
|
||||
if not sd.is_flux:
|
||||
raise ValueError("Only Flux models can load assistant adapters currently.")
|
||||
pipe = sd.pipeline
|
||||
print(f"Loading assistant adapter from {adapter_path}")
|
||||
adapter_name = adapter_path.split("/")[-1].split(".")[0]
|
||||
lora_state_dict = load_file(adapter_path)
|
||||
|
||||
linear_dim = int(lora_state_dict['transformer.single_transformer_blocks.0.attn.to_k.lora_A.weight'].shape[0])
|
||||
# linear_alpha = int(lora_state_dict['lora_transformer_single_transformer_blocks_0_attn_to_k.alpha'].item())
|
||||
linear_alpha = linear_dim
|
||||
transformer_only = 'transformer.proj_out.alpha' not in lora_state_dict
|
||||
# get dim and scale
|
||||
network_config = NetworkConfig(
|
||||
linear=linear_dim,
|
||||
linear_alpha=linear_alpha,
|
||||
transformer_only=transformer_only,
|
||||
)
|
||||
|
||||
network = LoRASpecialNetwork(
|
||||
text_encoder=pipe.text_encoder,
|
||||
unet=pipe.transformer,
|
||||
lora_dim=network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=network_config.linear_alpha,
|
||||
train_unet=True,
|
||||
train_text_encoder=False,
|
||||
is_flux=True,
|
||||
network_config=network_config,
|
||||
network_type=network_config.type,
|
||||
transformer_only=network_config.transformer_only,
|
||||
is_assistant_adapter=True
|
||||
)
|
||||
network.apply_to(
|
||||
pipe.text_encoder,
|
||||
pipe.transformer,
|
||||
apply_text_encoder=False,
|
||||
apply_unet=True
|
||||
)
|
||||
network.force_to(sd.device_torch, dtype=sd.torch_dtype)
|
||||
network.eval()
|
||||
network._update_torch_multiplier()
|
||||
network.load_weights(lora_state_dict)
|
||||
network.is_active = True
|
||||
|
||||
return network
|
||||
@@ -31,12 +31,18 @@ def get_mean_std(tensor):
|
||||
def adain(content_features, style_features):
|
||||
# Assumes that the content and style features are of shape (batch_size, channels, width, height)
|
||||
|
||||
dims = [2, 3]
|
||||
if len(content_features.shape) == 3:
|
||||
# content_features = content_features.unsqueeze(0)
|
||||
# style_features = style_features.unsqueeze(0)
|
||||
dims = [1]
|
||||
|
||||
# Step 1: Calculate mean and variance of content features
|
||||
content_mean, content_var = torch.mean(content_features, dim=[2, 3], keepdim=True), torch.var(content_features,
|
||||
dim=[2, 3],
|
||||
content_mean, content_var = torch.mean(content_features, dim=dims, keepdim=True), torch.var(content_features,
|
||||
dim=dims,
|
||||
keepdim=True)
|
||||
# Step 2: Calculate mean and variance of style features
|
||||
style_mean, style_var = torch.mean(style_features, dim=[2, 3], keepdim=True), torch.var(style_features, dim=[2, 3],
|
||||
style_mean, style_var = torch.mean(style_features, dim=dims, keepdim=True), torch.var(style_features, dim=dims,
|
||||
keepdim=True)
|
||||
|
||||
# Step 3: Normalize content features
|
||||
|
||||
@@ -51,6 +51,53 @@ resolutions_1024: List[BucketResolution] = [
|
||||
{"width": 512, "height": 1920},
|
||||
{"width": 512, "height": 1984},
|
||||
{"width": 512, "height": 2048},
|
||||
# extra wides
|
||||
{"width": 8192, "height": 128},
|
||||
{"width": 128, "height": 8192},
|
||||
]
|
||||
|
||||
# Even numbers so they can be patched easier
|
||||
resolutions_dit_1024: List[BucketResolution] = [
|
||||
# Base resolution
|
||||
{"width": 1024, "height": 1024},
|
||||
# widescreen
|
||||
{"width": 2048, "height": 512},
|
||||
{"width": 1792, "height": 576},
|
||||
{"width": 1728, "height": 576},
|
||||
{"width": 1664, "height": 576},
|
||||
{"width": 1600, "height": 640},
|
||||
{"width": 1536, "height": 640},
|
||||
{"width": 1472, "height": 704},
|
||||
{"width": 1408, "height": 704},
|
||||
{"width": 1344, "height": 704},
|
||||
{"width": 1344, "height": 768},
|
||||
{"width": 1280, "height": 768},
|
||||
{"width": 1216, "height": 832},
|
||||
{"width": 1152, "height": 832},
|
||||
{"width": 1152, "height": 896},
|
||||
{"width": 1088, "height": 896},
|
||||
{"width": 1088, "height": 960},
|
||||
{"width": 1024, "height": 960},
|
||||
# portrait
|
||||
{"width": 960, "height": 1024},
|
||||
{"width": 960, "height": 1088},
|
||||
{"width": 896, "height": 1088},
|
||||
{"width": 896, "height": 1152}, # 2:3
|
||||
{"width": 832, "height": 1152},
|
||||
{"width": 832, "height": 1216},
|
||||
{"width": 768, "height": 1280},
|
||||
{"width": 768, "height": 1344},
|
||||
{"width": 704, "height": 1408},
|
||||
{"width": 704, "height": 1472},
|
||||
{"width": 640, "height": 1536},
|
||||
{"width": 640, "height": 1600},
|
||||
{"width": 576, "height": 1664},
|
||||
{"width": 576, "height": 1728},
|
||||
{"width": 576, "height": 1792},
|
||||
{"width": 512, "height": 1856},
|
||||
{"width": 512, "height": 1920},
|
||||
{"width": 512, "height": 1984},
|
||||
{"width": 512, "height": 2048},
|
||||
]
|
||||
|
||||
|
||||
|
||||
406
toolkit/clip_vision_adapter.py
Normal file
406
toolkit/clip_vision_adapter.py
Normal file
@@ -0,0 +1,406 @@
|
||||
from typing import TYPE_CHECKING, Mapping, Any
|
||||
|
||||
import torch
|
||||
import weakref
|
||||
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPVisionModelWithProjection,
|
||||
CLIPVisionModel
|
||||
)
|
||||
|
||||
from toolkit.resampler import Resampler
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class Embedder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_input_tokens: int = 1,
|
||||
input_dim: int = 1024,
|
||||
num_output_tokens: int = 8,
|
||||
output_dim: int = 768,
|
||||
mid_dim: int = 1024
|
||||
):
|
||||
super(Embedder, self).__init__()
|
||||
self.num_output_tokens = num_output_tokens
|
||||
self.num_input_tokens = num_input_tokens
|
||||
self.input_dim = input_dim
|
||||
self.output_dim = output_dim
|
||||
|
||||
self.layer_norm = nn.LayerNorm(input_dim)
|
||||
self.fc1 = nn.Linear(input_dim, mid_dim)
|
||||
self.gelu = nn.GELU()
|
||||
# self.fc2 = nn.Linear(mid_dim, mid_dim)
|
||||
self.fc2 = nn.Linear(mid_dim, mid_dim)
|
||||
|
||||
self.fc2.weight.data.zero_()
|
||||
|
||||
self.layer_norm2 = nn.LayerNorm(mid_dim)
|
||||
self.fc3 = nn.Linear(mid_dim, mid_dim)
|
||||
self.gelu2 = nn.GELU()
|
||||
self.fc4 = nn.Linear(mid_dim, output_dim * num_output_tokens)
|
||||
|
||||
# set the weights to 0
|
||||
self.fc3.weight.data.zero_()
|
||||
self.fc4.weight.data.zero_()
|
||||
|
||||
|
||||
# self.static_tokens = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
|
||||
# self.scaler = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
|
||||
|
||||
def forward(self, x):
|
||||
if len(x.shape) == 2:
|
||||
x = x.unsqueeze(1)
|
||||
x = self.layer_norm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.gelu(x)
|
||||
x = self.fc2(x)
|
||||
x = self.layer_norm2(x)
|
||||
x = self.fc3(x)
|
||||
x = self.gelu2(x)
|
||||
x = self.fc4(x)
|
||||
|
||||
x = x.view(-1, self.num_output_tokens, self.output_dim)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ClipVisionAdapter(torch.nn.Module):
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: AdapterConfig):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.trigger = adapter_config.trigger
|
||||
self.trigger_class_name = adapter_config.trigger_class_name
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
# embedding stuff
|
||||
self.text_encoder_list = sd.text_encoder if isinstance(sd.text_encoder, list) else [sd.text_encoder]
|
||||
self.tokenizer_list = sd.tokenizer if isinstance(sd.tokenizer, list) else [sd.tokenizer]
|
||||
placeholder_tokens = [self.trigger]
|
||||
|
||||
# add dummy tokens for multi-vector
|
||||
additional_tokens = []
|
||||
for i in range(1, self.config.num_tokens):
|
||||
additional_tokens.append(f"{self.trigger}_{i}")
|
||||
placeholder_tokens += additional_tokens
|
||||
|
||||
# handle dual tokenizer
|
||||
self.tokenizer_list = self.sd_ref().tokenizer if isinstance(self.sd_ref().tokenizer, list) else [
|
||||
self.sd_ref().tokenizer]
|
||||
self.text_encoder_list = self.sd_ref().text_encoder if isinstance(self.sd_ref().text_encoder, list) else [
|
||||
self.sd_ref().text_encoder]
|
||||
|
||||
self.placeholder_token_ids = []
|
||||
self.embedding_tokens = []
|
||||
|
||||
print(f"Adding {placeholder_tokens} tokens to tokenizer")
|
||||
print(f"Adding {self.config.num_tokens} tokens to tokenizer")
|
||||
|
||||
|
||||
for text_encoder, tokenizer in zip(self.text_encoder_list, self.tokenizer_list):
|
||||
num_added_tokens = tokenizer.add_tokens(placeholder_tokens)
|
||||
if num_added_tokens != self.config.num_tokens:
|
||||
raise ValueError(
|
||||
f"The tokenizer already contains the token {self.trigger}. Please pass a different"
|
||||
f" `placeholder_token` that is not already in the tokenizer. Only added {num_added_tokens}"
|
||||
)
|
||||
|
||||
# Convert the initializer_token, placeholder_token to ids
|
||||
init_token_ids = tokenizer.encode(self.config.trigger_class_name, add_special_tokens=False)
|
||||
# if length of token ids is more than number of orm embedding tokens fill with *
|
||||
if len(init_token_ids) > self.config.num_tokens:
|
||||
init_token_ids = init_token_ids[:self.config.num_tokens]
|
||||
elif len(init_token_ids) < self.config.num_tokens:
|
||||
pad_token_id = tokenizer.encode(["*"], add_special_tokens=False)
|
||||
init_token_ids += pad_token_id * (self.config.num_tokens - len(init_token_ids))
|
||||
|
||||
placeholder_token_ids = tokenizer.encode(placeholder_tokens, add_special_tokens=False)
|
||||
self.placeholder_token_ids.append(placeholder_token_ids)
|
||||
|
||||
# Resize the token embeddings as we are adding new special tokens to the tokenizer
|
||||
text_encoder.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
# Initialise the newly added placeholder token with the embeddings of the initializer token
|
||||
token_embeds = text_encoder.get_input_embeddings().weight.data
|
||||
with torch.no_grad():
|
||||
for initializer_token_id, token_id in zip(init_token_ids, placeholder_token_ids):
|
||||
token_embeds[token_id] = token_embeds[initializer_token_id].clone()
|
||||
|
||||
# replace "[name] with this. on training. This is automatically generated in pipeline on inference
|
||||
self.embedding_tokens.append(" ".join(tokenizer.convert_ids_to_tokens(placeholder_token_ids)))
|
||||
|
||||
# backup text encoder embeddings
|
||||
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
|
||||
|
||||
try:
|
||||
self.clip_image_processor = CLIPImageProcessor.from_pretrained(self.config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.clip_image_processor = CLIPImageProcessor()
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
self.config.image_encoder_path,
|
||||
ignore_mismatched_sizes=True
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
if self.config.train_image_encoder:
|
||||
self.image_encoder.train()
|
||||
else:
|
||||
self.image_encoder.eval()
|
||||
|
||||
# max_seq_len = CLIP tokens + CLS token
|
||||
image_encoder_state_dict = self.image_encoder.state_dict()
|
||||
in_tokens = 257
|
||||
if "vision_model.embeddings.position_embedding.weight" in image_encoder_state_dict:
|
||||
# clip
|
||||
in_tokens = int(image_encoder_state_dict["vision_model.embeddings.position_embedding.weight"].shape[0])
|
||||
|
||||
if hasattr(self.image_encoder.config, 'hidden_sizes'):
|
||||
embedding_dim = self.image_encoder.config.hidden_sizes[-1]
|
||||
else:
|
||||
embedding_dim = self.image_encoder.config.target_hidden_size
|
||||
|
||||
if self.config.clip_layer == 'image_embeds':
|
||||
in_tokens = 1
|
||||
embedding_dim = self.image_encoder.config.projection_dim
|
||||
|
||||
self.embedder = Embedder(
|
||||
num_output_tokens=self.config.num_tokens,
|
||||
num_input_tokens=in_tokens,
|
||||
input_dim=embedding_dim,
|
||||
output_dim=self.sd_ref().unet.config['cross_attention_dim'],
|
||||
mid_dim=embedding_dim * self.config.num_tokens,
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
|
||||
self.embedder.train()
|
||||
|
||||
def state_dict(self, *args, destination=None, prefix='', keep_vars=False):
|
||||
state_dict = {
|
||||
'embedder': self.embedder.state_dict(*args, destination=destination, prefix=prefix, keep_vars=keep_vars)
|
||||
}
|
||||
if self.config.train_image_encoder:
|
||||
state_dict['image_encoder'] = self.image_encoder.state_dict(
|
||||
*args, destination=destination, prefix=prefix,
|
||||
keep_vars=keep_vars)
|
||||
|
||||
return state_dict
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||
self.embedder.load_state_dict(state_dict["embedder"], strict=strict)
|
||||
if self.config.train_image_encoder and 'image_encoder' in state_dict:
|
||||
self.image_encoder.load_state_dict(state_dict["image_encoder"], strict=strict)
|
||||
|
||||
def parameters(self, *args, **kwargs):
|
||||
yield from self.embedder.parameters(*args, **kwargs)
|
||||
|
||||
def named_parameters(self, *args, **kwargs):
|
||||
yield from self.embedder.named_parameters(*args, **kwargs)
|
||||
|
||||
def get_clip_image_embeds_from_tensors(
|
||||
self, tensors_0_1: torch.Tensor, drop=False,
|
||||
is_training=False,
|
||||
has_been_preprocessed=False
|
||||
) -> torch.Tensor:
|
||||
with torch.no_grad():
|
||||
if not has_been_preprocessed:
|
||||
# tensors should be 0-1
|
||||
if tensors_0_1.ndim == 3:
|
||||
tensors_0_1 = tensors_0_1.unsqueeze(0)
|
||||
# training tensors are 0 - 1
|
||||
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
|
||||
|
||||
# if images are out of this range throw error
|
||||
if tensors_0_1.min() < -0.3 or tensors_0_1.max() > 1.3:
|
||||
raise ValueError("image tensor values must be between 0 and 1. Got min: {}, max: {}".format(
|
||||
tensors_0_1.min(), tensors_0_1.max()
|
||||
))
|
||||
# unconditional
|
||||
if drop:
|
||||
if self.clip_noise_zero:
|
||||
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
|
||||
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
|
||||
dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
tensors_0_1 = tensors_0_1 * noise_scale
|
||||
else:
|
||||
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
|
||||
# tensors_0_1 = tensors_0_1 * 0
|
||||
clip_image = self.clip_image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
else:
|
||||
if drop:
|
||||
# scale the noise down
|
||||
if self.clip_noise_zero:
|
||||
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
|
||||
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
|
||||
dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
tensors_0_1 = tensors_0_1 * noise_scale
|
||||
else:
|
||||
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
|
||||
# tensors_0_1 = tensors_0_1 * 0
|
||||
mean = torch.tensor(self.clip_image_processor.image_mean).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
std = torch.tensor(self.clip_image_processor.image_std).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
|
||||
clip_image = (tensors_0_1 - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
|
||||
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
clip_image = clip_image.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype)).detach()
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training:
|
||||
self.image_encoder.train()
|
||||
else:
|
||||
self.image_encoder.eval()
|
||||
clip_output = self.image_encoder(clip_image, output_hidden_states=True)
|
||||
|
||||
if self.config.clip_layer == 'penultimate_hidden_states':
|
||||
# they skip last layer for ip+
|
||||
# https://github.com/tencent-ailab/IP-Adapter/blob/f4b6742db35ea6d81c7b829a55b0a312c7f5a677/tutorial_train_plus.py#L403C26-L403C26
|
||||
clip_image_embeds = clip_output.hidden_states[-2]
|
||||
elif self.config.clip_layer == 'last_hidden_state':
|
||||
clip_image_embeds = clip_output.hidden_states[-1]
|
||||
else:
|
||||
clip_image_embeds = clip_output.image_embeds
|
||||
return clip_image_embeds
|
||||
|
||||
import torch
|
||||
|
||||
def set_vec(self, new_vector, text_encoder_idx=0):
|
||||
# Get the embedding layer
|
||||
embedding_layer = self.text_encoder_list[text_encoder_idx].get_input_embeddings()
|
||||
|
||||
# Indices to replace in the embeddings
|
||||
indices_to_replace = self.placeholder_token_ids[text_encoder_idx]
|
||||
|
||||
# Replace the specified embeddings with new_vector
|
||||
for idx in indices_to_replace:
|
||||
vector_idx = idx - indices_to_replace[0]
|
||||
embedding_layer.weight[idx] = new_vector[vector_idx]
|
||||
|
||||
# adds it to the tokenizer
|
||||
def forward(self, clip_image_embeds: torch.Tensor) -> PromptEmbeds:
|
||||
clip_image_embeds = clip_image_embeds.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
if clip_image_embeds.ndim == 2:
|
||||
# expand the token dimension
|
||||
clip_image_embeds = clip_image_embeds.unsqueeze(1)
|
||||
image_prompt_embeds = self.embedder(clip_image_embeds)
|
||||
# todo add support for multiple batch sizes
|
||||
if image_prompt_embeds.shape[0] != 1:
|
||||
raise ValueError("Batch size must be 1 for embedder for now")
|
||||
|
||||
# output on sd1.5 is bs, num_tokens, 768
|
||||
if len(self.text_encoder_list) == 1:
|
||||
# add it to the text encoder
|
||||
self.set_vec(image_prompt_embeds[0], text_encoder_idx=0)
|
||||
elif len(self.text_encoder_list) == 2:
|
||||
if self.text_encoder_list[0].config.target_hidden_size + self.text_encoder_list[1].config.target_hidden_size != \
|
||||
image_prompt_embeds.shape[2]:
|
||||
raise ValueError("Something went wrong. The embeddings do not match the text encoder sizes")
|
||||
# sdxl variants
|
||||
# image_prompt_embeds = 2048
|
||||
# te1 = 768
|
||||
# te2 = 1280
|
||||
te1_embeds = image_prompt_embeds[:, :, :self.text_encoder_list[0].config.target_hidden_size]
|
||||
te2_embeds = image_prompt_embeds[:, :, self.text_encoder_list[0].config.target_hidden_size:]
|
||||
self.set_vec(te1_embeds[0], text_encoder_idx=0)
|
||||
self.set_vec(te2_embeds[0], text_encoder_idx=1)
|
||||
else:
|
||||
|
||||
raise ValueError("Unsupported number of text encoders")
|
||||
# just a place to put a breakpoint
|
||||
pass
|
||||
|
||||
def restore_embeddings(self):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(
|
||||
self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids
|
||||
):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[
|
||||
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
with torch.no_grad():
|
||||
text_encoder.get_input_embeddings().weight[
|
||||
index_no_updates
|
||||
] = orig_embeds[index_no_updates]
|
||||
# detach it all
|
||||
text_encoder.get_input_embeddings().weight.detach_()
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.image_encoder.gradient_checkpointing = True
|
||||
|
||||
def inject_trigger_into_prompt(self, prompt, expand_token=False, to_replace_list=None, add_if_not_present=True):
|
||||
output_prompt = prompt
|
||||
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
|
||||
default_replacements = ["[name]", "[trigger]"]
|
||||
|
||||
replace_with = embedding_tokens if expand_token else self.trigger
|
||||
if to_replace_list is None:
|
||||
to_replace_list = default_replacements
|
||||
else:
|
||||
to_replace_list += default_replacements
|
||||
|
||||
# remove duplicates
|
||||
to_replace_list = list(set(to_replace_list))
|
||||
|
||||
# replace them all
|
||||
for to_replace in to_replace_list:
|
||||
# replace it
|
||||
output_prompt = output_prompt.replace(to_replace, replace_with)
|
||||
|
||||
# see how many times replace_with is in the prompt
|
||||
num_instances = output_prompt.count(replace_with)
|
||||
|
||||
if num_instances == 0 and add_if_not_present:
|
||||
# add it to the beginning of the prompt
|
||||
output_prompt = replace_with + " " + output_prompt
|
||||
|
||||
if num_instances > 1:
|
||||
print(
|
||||
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
|
||||
|
||||
return output_prompt
|
||||
|
||||
# reverses injection with class name. useful for normalizations
|
||||
def inject_trigger_class_name_into_prompt(self, prompt):
|
||||
output_prompt = prompt
|
||||
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
|
||||
|
||||
default_replacements = ["[name]", "[trigger]", embedding_tokens, self.trigger]
|
||||
|
||||
replace_with = self.config.trigger_class_name
|
||||
to_replace_list = default_replacements
|
||||
|
||||
# remove duplicates
|
||||
to_replace_list = list(set(to_replace_list))
|
||||
|
||||
# replace them all
|
||||
for to_replace in to_replace_list:
|
||||
# replace it
|
||||
output_prompt = output_prompt.replace(to_replace, replace_with)
|
||||
|
||||
# see how many times replace_with is in the prompt
|
||||
num_instances = output_prompt.count(replace_with)
|
||||
|
||||
if num_instances > 1:
|
||||
print(
|
||||
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
|
||||
|
||||
return output_prompt
|
||||
@@ -43,9 +43,7 @@ def preprocess_config(config: OrderedDict, name: str = None):
|
||||
if "name" not in config["config"] and name is None:
|
||||
raise ValueError("config file must have a config.name key")
|
||||
# we need to replace tags. For now just [name]
|
||||
if name is not None:
|
||||
config["config"]["name"] = name
|
||||
else:
|
||||
if name is None:
|
||||
name = config["config"]["name"]
|
||||
config_string = json.dumps(config)
|
||||
config_string = config_string.replace("[name]", name)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
import time
|
||||
from typing import List, Optional, Literal, Union
|
||||
from typing import List, Optional, Literal, Union, TYPE_CHECKING, Dict
|
||||
import random
|
||||
|
||||
import torch
|
||||
@@ -11,16 +11,21 @@ ImgExt = Literal['jpg', 'png', 'webp']
|
||||
|
||||
SaveFormat = Literal['safetensors', 'diffusers']
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.guidance import GuidanceType
|
||||
|
||||
|
||||
class SaveConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.save_every: int = kwargs.get('save_every', 1000)
|
||||
self.dtype: str = kwargs.get('save_dtype', 'float16')
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.max_step_saves_to_keep: int = kwargs.get('max_step_saves_to_keep', 5)
|
||||
self.save_format: SaveFormat = kwargs.get('save_format', 'safetensors')
|
||||
if self.save_format not in ['safetensors', 'diffusers']:
|
||||
raise ValueError(f"save_format must be safetensors or diffusers, got {self.save_format}")
|
||||
|
||||
self.push_to_hub: bool = kwargs.get("push_to_hub", False)
|
||||
self.hf_repo_id: Optional[str] = kwargs.get("hf_repo_id", None)
|
||||
self.hf_private: Optional[str] = kwargs.get("hf_private", False)
|
||||
|
||||
class LogingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -45,7 +50,9 @@ class SampleConfig:
|
||||
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
|
||||
self.ext: ImgExt = kwargs.get('format', 'jpg')
|
||||
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5) # step to start using refiner on sample if it exists
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at',
|
||||
0.5) # step to start using refiner on sample if it exists
|
||||
self.extra_values = kwargs.get('extra_values', [])
|
||||
|
||||
|
||||
class LormModuleSettingsConfig:
|
||||
@@ -109,6 +116,7 @@ class NetworkConfig:
|
||||
self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha)
|
||||
self.conv_alpha: float = kwargs.get('conv_alpha', self.conv)
|
||||
self.dropout: Union[float, None] = kwargs.get('dropout', None)
|
||||
self.network_kwargs: dict = kwargs.get('network_kwargs', {})
|
||||
|
||||
self.lorm_config: Union[LoRMConfig, None] = None
|
||||
lorm = kwargs.get('lorm', None)
|
||||
@@ -122,13 +130,17 @@ class NetworkConfig:
|
||||
if self.lorm_config.do_conv:
|
||||
self.conv = 4
|
||||
|
||||
self.transformer_only = kwargs.get('transformer_only', True)
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+']
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net']
|
||||
|
||||
CLIPLayer = Literal['penultimate_hidden_states', 'image_embeds', 'last_hidden_state']
|
||||
|
||||
|
||||
class AdapterConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip
|
||||
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip, clip, control_net
|
||||
self.in_channels: int = kwargs.get('in_channels', 3)
|
||||
self.channels: List[int] = kwargs.get('channels', [320, 640, 1280, 1280])
|
||||
self.num_res_blocks: int = kwargs.get('num_res_blocks', 2)
|
||||
@@ -140,6 +152,57 @@ class AdapterConfig:
|
||||
self.image_encoder_path: str = kwargs.get('image_encoder_path', None)
|
||||
self.name_or_path = kwargs.get('name_or_path', None)
|
||||
|
||||
num_tokens = kwargs.get('num_tokens', None)
|
||||
if num_tokens is None and self.type.startswith('ip'):
|
||||
if self.type == 'ip+':
|
||||
num_tokens = 16
|
||||
num_tokens = 16
|
||||
elif self.type == 'ip':
|
||||
num_tokens = 4
|
||||
|
||||
self.num_tokens: int = num_tokens
|
||||
self.train_image_encoder: bool = kwargs.get('train_image_encoder', False)
|
||||
self.train_only_image_encoder: bool = kwargs.get('train_only_image_encoder', False)
|
||||
if self.train_only_image_encoder:
|
||||
self.train_image_encoder = True
|
||||
self.train_only_image_encoder_positional_embedding: bool = kwargs.get(
|
||||
'train_only_image_encoder_positional_embedding', False)
|
||||
self.image_encoder_arch: str = kwargs.get('image_encoder_arch', 'clip') # clip vit vit_hybrid, safe
|
||||
self.safe_reducer_channels: int = kwargs.get('safe_reducer_channels', 512)
|
||||
self.safe_channels: int = kwargs.get('safe_channels', 2048)
|
||||
self.safe_tokens: int = kwargs.get('safe_tokens', 8)
|
||||
self.quad_image: bool = kwargs.get('quad_image', False)
|
||||
|
||||
# clip vision
|
||||
self.trigger = kwargs.get('trigger', 'tri993r')
|
||||
self.trigger_class_name = kwargs.get('trigger_class_name', None)
|
||||
|
||||
self.class_names = kwargs.get('class_names', [])
|
||||
|
||||
self.clip_layer: CLIPLayer = kwargs.get('clip_layer', None)
|
||||
if self.clip_layer is None:
|
||||
if self.type.startswith('ip+'):
|
||||
self.clip_layer = 'penultimate_hidden_states'
|
||||
else:
|
||||
self.clip_layer = 'last_hidden_state'
|
||||
|
||||
# text encoder
|
||||
self.text_encoder_path: str = kwargs.get('text_encoder_path', None)
|
||||
self.text_encoder_arch: str = kwargs.get('text_encoder_arch', 'clip') # clip t5
|
||||
|
||||
self.train_scaler: bool = kwargs.get('train_scaler', False)
|
||||
self.scaler_lr: Optional[float] = kwargs.get('scaler_lr', None)
|
||||
|
||||
# trains with a scaler to easy channel bias but merges it in on save
|
||||
self.merge_scaler: bool = kwargs.get('merge_scaler', False)
|
||||
|
||||
# for ilora
|
||||
self.head_dim: int = kwargs.get('head_dim', 1024)
|
||||
self.num_heads: int = kwargs.get('num_heads', 1)
|
||||
self.ilora_down: bool = kwargs.get('ilora_down', True)
|
||||
self.ilora_mid: bool = kwargs.get('ilora_mid', True)
|
||||
self.ilora_up: bool = kwargs.get('ilora_up', True)
|
||||
|
||||
|
||||
class EmbeddingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -147,6 +210,7 @@ class EmbeddingConfig:
|
||||
self.tokens = kwargs.get('tokens', 4)
|
||||
self.init_words = kwargs.get('init_words', '*')
|
||||
self.save_format = kwargs.get('save_format', 'safetensors')
|
||||
self.trigger_class_name = kwargs.get('trigger_class_name', None) # used for inverted masked prior
|
||||
|
||||
|
||||
ContentOrStyleType = Literal['balanced', 'style', 'content']
|
||||
@@ -157,6 +221,7 @@ class TrainConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.noise_scheduler = kwargs.get('noise_scheduler', 'ddpm')
|
||||
self.content_or_style: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
|
||||
self.content_or_style_reg: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
|
||||
self.steps: int = kwargs.get('steps', 1000)
|
||||
self.lr = kwargs.get('lr', 1e-6)
|
||||
self.unet_lr = kwargs.get('unet_lr', self.lr)
|
||||
@@ -175,8 +240,10 @@ class TrainConfig:
|
||||
self.xformers = kwargs.get('xformers', False)
|
||||
self.sdp = kwargs.get('sdp', False)
|
||||
self.train_unet = kwargs.get('train_unet', True)
|
||||
self.train_text_encoder = kwargs.get('train_text_encoder', True)
|
||||
self.train_text_encoder = kwargs.get('train_text_encoder', False)
|
||||
self.train_refiner = kwargs.get('train_refiner', True)
|
||||
self.train_turbo = kwargs.get('train_turbo', False)
|
||||
self.show_turbo_outputs = kwargs.get('show_turbo_outputs', False)
|
||||
self.min_snr_gamma = kwargs.get('min_snr_gamma', None)
|
||||
self.snr_gamma = kwargs.get('snr_gamma', None)
|
||||
# trains a gamma, offset, and scale to adjust loss to adapt to timestep differentials
|
||||
@@ -184,6 +251,7 @@ class TrainConfig:
|
||||
self.learnable_snr_gos = kwargs.get('learnable_snr_gos', False)
|
||||
self.noise_offset = kwargs.get('noise_offset', 0.0)
|
||||
self.skip_first_sample = kwargs.get('skip_first_sample', False)
|
||||
self.force_first_sample = kwargs.get('force_first_sample', False)
|
||||
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True)
|
||||
self.weight_jitter = kwargs.get('weight_jitter', 0.0)
|
||||
self.merge_network_on_save = kwargs.get('merge_network_on_save', False)
|
||||
@@ -191,12 +259,20 @@ class TrainConfig:
|
||||
self.start_step = kwargs.get('start_step', None)
|
||||
self.free_u = kwargs.get('free_u', False)
|
||||
self.adapter_assist_name_or_path: Optional[str] = kwargs.get('adapter_assist_name_or_path', None)
|
||||
self.adapter_assist_type: Optional[str] = kwargs.get('adapter_assist_type', 't2i') # t2i, control_net
|
||||
self.noise_multiplier = kwargs.get('noise_multiplier', 1.0)
|
||||
self.target_noise_multiplier = kwargs.get('target_noise_multiplier', 1.0)
|
||||
self.img_multiplier = kwargs.get('img_multiplier', 1.0)
|
||||
self.noisy_latent_multiplier = kwargs.get('noisy_latent_multiplier', 1.0)
|
||||
self.latent_multiplier = kwargs.get('latent_multiplier', 1.0)
|
||||
self.negative_prompt = kwargs.get('negative_prompt', None)
|
||||
self.max_negative_prompts = kwargs.get('max_negative_prompts', 1)
|
||||
# multiplier applied to loos on regularization images
|
||||
self.reg_weight = kwargs.get('reg_weight', 1.0)
|
||||
self.num_train_timesteps = kwargs.get('num_train_timesteps', 1000)
|
||||
self.random_noise_shift = kwargs.get('random_noise_shift', 0.0)
|
||||
# automatically adapte the vae scaling based on the image norm
|
||||
self.adaptive_scaling_factor = kwargs.get('adaptive_scaling_factor', False)
|
||||
|
||||
# dropout that happens before encoding. It functions independently per text encoder
|
||||
self.prompt_dropout_prob = kwargs.get('prompt_dropout_prob', 0.0)
|
||||
@@ -238,19 +314,79 @@ class TrainConfig:
|
||||
if match_adapter_assist and self.match_adapter_chance == 0.0:
|
||||
self.match_adapter_chance = 1.0
|
||||
|
||||
# standardize inputs to the meand std of the model knowledge
|
||||
self.standardize_images = kwargs.get('standardize_images', False)
|
||||
self.standardize_latents = kwargs.get('standardize_latents', False)
|
||||
|
||||
if self.train_turbo and not self.noise_scheduler.startswith("euler"):
|
||||
raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers")
|
||||
|
||||
self.dynamic_noise_offset = kwargs.get('dynamic_noise_offset', False)
|
||||
self.do_cfg = kwargs.get('do_cfg', False)
|
||||
self.do_random_cfg = kwargs.get('do_random_cfg', False)
|
||||
self.cfg_scale = kwargs.get('cfg_scale', 1.0)
|
||||
self.max_cfg_scale = kwargs.get('max_cfg_scale', self.cfg_scale)
|
||||
self.cfg_rescale = kwargs.get('cfg_rescale', None)
|
||||
if self.cfg_rescale is None:
|
||||
self.cfg_rescale = self.cfg_scale
|
||||
|
||||
# applies the inverse of the prediction mean and std to the target to correct
|
||||
# for norm drift
|
||||
self.correct_pred_norm = kwargs.get('correct_pred_norm', False)
|
||||
self.correct_pred_norm_multiplier = kwargs.get('correct_pred_norm_multiplier', 1.0)
|
||||
|
||||
self.loss_type = kwargs.get('loss_type', 'mse')
|
||||
|
||||
# scale the prediction by this. Increase for more detail, decrease for less
|
||||
self.pred_scaler = kwargs.get('pred_scaler', 1.0)
|
||||
|
||||
# repeats the prompt a few times to saturate the encoder
|
||||
self.prompt_saturation_chance = kwargs.get('prompt_saturation_chance', 0.0)
|
||||
|
||||
# applies negative loss on the prior to encourage network to diverge from it
|
||||
self.do_prior_divergence = kwargs.get('do_prior_divergence', False)
|
||||
|
||||
ema_config: Union[Dict, None] = kwargs.get('ema_config', None)
|
||||
if ema_config is not None:
|
||||
ema_config['use_ema'] = True
|
||||
print(f"Using EMA")
|
||||
else:
|
||||
ema_config = {'use_ema': False}
|
||||
|
||||
self.ema_config: EMAConfig = EMAConfig(**ema_config)
|
||||
|
||||
# adds an additional loss to the network to encourage it output a normalized standard deviation
|
||||
self.target_norm_std = kwargs.get('target_norm_std', None)
|
||||
self.target_norm_std_value = kwargs.get('target_norm_std_value', 1.0)
|
||||
self.linear_timesteps = kwargs.get('linear_timesteps', False)
|
||||
self.disable_sampling = kwargs.get('disable_sampling', False)
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.name_or_path: str = kwargs.get('name_or_path', None)
|
||||
self.is_v2: bool = kwargs.get('is_v2', False)
|
||||
self.is_xl: bool = kwargs.get('is_xl', False)
|
||||
self.is_pixart: bool = kwargs.get('is_pixart', False)
|
||||
self.is_pixart_sigma: bool = kwargs.get('is_pixart_sigma', False)
|
||||
self.is_auraflow: bool = kwargs.get('is_auraflow', False)
|
||||
self.is_v3: bool = kwargs.get('is_v3', False)
|
||||
self.is_flux: bool = kwargs.get('is_flux', False)
|
||||
if self.is_pixart_sigma:
|
||||
self.is_pixart = True
|
||||
self.use_flux_cfg = kwargs.get('use_flux_cfg', False)
|
||||
self.is_ssd: bool = kwargs.get('is_ssd', False)
|
||||
self.is_vega: bool = kwargs.get('is_vega', False)
|
||||
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
|
||||
self.dtype: str = kwargs.get('dtype', 'float16')
|
||||
self.vae_path = kwargs.get('vae_path', None)
|
||||
self.refiner_name_or_path = kwargs.get('refiner_name_or_path', None)
|
||||
self._original_refiner_name_or_path = self.refiner_name_or_path
|
||||
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5)
|
||||
self.lora_path = kwargs.get('lora_path', None)
|
||||
# mainly for decompression loras for distilled models
|
||||
self.assistant_lora_path = kwargs.get('assistant_lora_path', None)
|
||||
self.latent_space_version = kwargs.get('latent_space_version', None)
|
||||
|
||||
# only for SDXL models for now
|
||||
self.use_text_encoder_1: bool = kwargs.get('use_text_encoder_1', True)
|
||||
@@ -265,6 +401,31 @@ class ModelConfig:
|
||||
# sed sdxl as true since it is mostly the same architecture
|
||||
self.is_xl = True
|
||||
|
||||
if self.is_vega:
|
||||
self.is_xl = True
|
||||
|
||||
# for text encoder quant. Only works with pixart currently
|
||||
self.text_encoder_bits = kwargs.get('text_encoder_bits', 16) # 16, 8, 4
|
||||
self.unet_path = kwargs.get("unet_path", None)
|
||||
self.unet_sample_size = kwargs.get("unet_sample_size", None)
|
||||
self.vae_device = kwargs.get("vae_device", None)
|
||||
self.vae_dtype = kwargs.get("vae_dtype", self.dtype)
|
||||
self.te_device = kwargs.get("te_device", None)
|
||||
self.te_dtype = kwargs.get("te_dtype", self.dtype)
|
||||
|
||||
# only for flux for now
|
||||
self.quantize = kwargs.get("quantize", False)
|
||||
self.low_vram = kwargs.get("low_vram", False)
|
||||
pass
|
||||
|
||||
|
||||
class EMAConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.use_ema: bool = kwargs.get('use_ema', False)
|
||||
self.ema_decay: float = kwargs.get('ema_decay', 0.999)
|
||||
# feeds back the decay difference into the parameter
|
||||
self.use_feedback: bool = kwargs.get('use_feedback', False)
|
||||
|
||||
|
||||
class ReferenceDatasetConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -352,29 +513,41 @@ class DatasetConfig:
|
||||
self.dataset_path: str = kwargs.get('dataset_path', None)
|
||||
|
||||
self.default_caption: str = kwargs.get('default_caption', None)
|
||||
self.random_triggers: List[str] = kwargs.get('random_triggers', [])
|
||||
random_triggers = kwargs.get('random_triggers', [])
|
||||
# if they are a string, load them from a file
|
||||
if isinstance(random_triggers, str) and os.path.exists(random_triggers):
|
||||
with open(random_triggers, 'r') as f:
|
||||
random_triggers = f.read().splitlines()
|
||||
# remove empty lines
|
||||
random_triggers = [line for line in random_triggers if line.strip() != '']
|
||||
self.random_triggers: List[str] = random_triggers
|
||||
self.random_triggers_max: int = kwargs.get('random_triggers_max', 1)
|
||||
self.caption_ext: str = kwargs.get('caption_ext', None)
|
||||
self.random_scale: bool = kwargs.get('random_scale', False)
|
||||
self.random_crop: bool = kwargs.get('random_crop', False)
|
||||
self.resolution: int = kwargs.get('resolution', 512)
|
||||
self.scale: float = kwargs.get('scale', 1.0)
|
||||
self.buckets: bool = kwargs.get('buckets', False)
|
||||
self.buckets: bool = kwargs.get('buckets', True)
|
||||
self.bucket_tolerance: int = kwargs.get('bucket_tolerance', 64)
|
||||
self.is_reg: bool = kwargs.get('is_reg', False)
|
||||
self.network_weight: float = float(kwargs.get('network_weight', 1.0))
|
||||
self.token_dropout_rate: float = float(kwargs.get('token_dropout_rate', 0.0))
|
||||
self.shuffle_tokens: bool = kwargs.get('shuffle_tokens', False)
|
||||
self.caption_dropout_rate: float = float(kwargs.get('caption_dropout_rate', 0.0))
|
||||
self.keep_tokens: int = kwargs.get('keep_tokens', 0) # #of first tokens to always keep unless caption dropped
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_y', False)
|
||||
self.augments: List[str] = kwargs.get('augments', [])
|
||||
self.control_path: str = kwargs.get('control_path', None) # depth maps, etc
|
||||
# instead of cropping ot match image, it will serve the full size control image (clip images ie for ip adapters)
|
||||
self.full_size_control_images: bool = kwargs.get('full_size_control_images', False)
|
||||
self.alpha_mask: bool = kwargs.get('alpha_mask', False) # if true, will use alpha channel as mask
|
||||
self.mask_path: str = kwargs.get('mask_path',
|
||||
None) # focus mask (black and white. White has higher loss than black)
|
||||
self.unconditional_path: str = kwargs.get('unconditional_path', None) # path where matching unconditional images are located
|
||||
self.unconditional_path: str = kwargs.get('unconditional_path',
|
||||
None) # path where matching unconditional images are located
|
||||
self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask
|
||||
self.mask_min_value: float = kwargs.get('mask_min_value', 0.01) # min value for . 0 - 1
|
||||
self.mask_min_value: float = kwargs.get('mask_min_value', 0.0) # min value for . 0 - 1
|
||||
self.poi: Union[str, None] = kwargs.get('poi',
|
||||
None) # if one is set and in json data, will be used as auto crop scale point of interes
|
||||
self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset
|
||||
@@ -382,6 +555,9 @@ class DatasetConfig:
|
||||
self.cache_latents: bool = kwargs.get('cache_latents', False)
|
||||
# cache latents to disk will store them on disk. If both are true, it will save to disk, but keep in memory
|
||||
self.cache_latents_to_disk: bool = kwargs.get('cache_latents_to_disk', False)
|
||||
self.cache_clip_vision_to_disk: bool = kwargs.get('cache_clip_vision_to_disk', False)
|
||||
|
||||
self.standardize_images: bool = kwargs.get('standardize_images', False)
|
||||
|
||||
# https://albumentations.ai/docs/api_reference/augmentations/transforms
|
||||
# augmentations are returned as a separate image and cannot currently be cached
|
||||
@@ -400,6 +576,21 @@ class DatasetConfig:
|
||||
if legacy_caption_type:
|
||||
self.caption_ext = legacy_caption_type
|
||||
self.caption_type = self.caption_ext
|
||||
self.guidance_type: GuidanceType = kwargs.get('guidance_type', 'targeted')
|
||||
|
||||
# ip adapter / reference dataset
|
||||
self.clip_image_path: str = kwargs.get('clip_image_path', None) # depth maps, etc
|
||||
self.clip_image_augmentations: List[dict] = kwargs.get('clip_image_augmentations', None)
|
||||
self.clip_image_shuffle_augmentations: bool = kwargs.get('clip_image_shuffle_augmentations', False)
|
||||
self.replacements: List[str] = kwargs.get('replacements', [])
|
||||
self.loss_multiplier: float = kwargs.get('loss_multiplier', 1.0)
|
||||
|
||||
self.num_workers: int = kwargs.get('num_workers', 2)
|
||||
self.prefetch_factor: int = kwargs.get('prefetch_factor', 2)
|
||||
self.extra_values: List[float] = kwargs.get('extra_values', [])
|
||||
self.square_crop: bool = kwargs.get('square_crop', False)
|
||||
# apply same augmentations to control images. Usually want this true unless special case
|
||||
self.replay_transforms: bool = kwargs.get('replay_transforms', True)
|
||||
|
||||
|
||||
def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]:
|
||||
@@ -448,6 +639,7 @@ class GenerateImageConfig:
|
||||
latents: Union[torch.Tensor | None] = None, # input latent to start with,
|
||||
extra_kwargs: dict = None, # extra data to save with prompt file
|
||||
refiner_start_at: float = 0.5, # start at this percentage of a step. 0.0 to 1.0 . 1.0 is the end
|
||||
extra_values: List[float] = None, # extra values to save with prompt file
|
||||
):
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
@@ -475,6 +667,7 @@ class GenerateImageConfig:
|
||||
self.adapter_conditioning_scale: float = adapter_conditioning_scale
|
||||
self.extra_kwargs = extra_kwargs if extra_kwargs is not None else {}
|
||||
self.refiner_start_at = refiner_start_at
|
||||
self.extra_values = extra_values if extra_values is not None else []
|
||||
|
||||
# prompt string will override any settings above
|
||||
self._process_prompt_string()
|
||||
@@ -484,7 +677,7 @@ class GenerateImageConfig:
|
||||
self.negative_prompt_2 = negative_prompt
|
||||
|
||||
if prompt_2 is None:
|
||||
self.prompt_2 = prompt
|
||||
self.prompt_2 = self.prompt
|
||||
|
||||
# parse prompt paths
|
||||
if self.output_path is None and self.output_folder is None:
|
||||
@@ -633,6 +826,12 @@ class GenerateImageConfig:
|
||||
self.adapter_conditioning_scale = float(content)
|
||||
elif flag == 'ref':
|
||||
self.refiner_start_at = float(content)
|
||||
elif flag == 'ev':
|
||||
# split by comma
|
||||
self.extra_values = [float(val) for val in content.split(',')]
|
||||
elif flag == 'extra_values':
|
||||
# split by comma
|
||||
self.extra_values = [float(val) for val in content.split(',')]
|
||||
|
||||
def post_process_embeddings(
|
||||
self,
|
||||
|
||||
911
toolkit/custom_adapter.py
Normal file
911
toolkit/custom_adapter.py
Normal file
@@ -0,0 +1,911 @@
|
||||
import torch
|
||||
import sys
|
||||
|
||||
from PIL import Image
|
||||
from torch.nn import Parameter
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection, T5EncoderModel, CLIPTextModel, \
|
||||
CLIPTokenizer, T5Tokenizer
|
||||
|
||||
from toolkit.models.clip_fusion import CLIPFusionModule
|
||||
from toolkit.models.clip_pre_processor import CLIPImagePreProcessor
|
||||
from toolkit.models.ilora import InstantLoRAModule
|
||||
from toolkit.models.single_value_adapter import SingleValueAdapter
|
||||
from toolkit.models.te_adapter import TEAdapter
|
||||
from toolkit.models.te_aug_adapter import TEAugAdapter
|
||||
from toolkit.models.vd_adapter import VisionDirectAdapter
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.photomaker import PhotoMakerIDEncoder, FuseModule, PhotoMakerCLIPEncoder
|
||||
from toolkit.saving import load_ip_adapter_model, load_custom_adapter_model
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from typing import TYPE_CHECKING, Union, Iterator, Mapping, Any, Tuple, List, Optional, Dict
|
||||
from collections import OrderedDict
|
||||
from ipadapter.ip_adapter.attention_processor import AttnProcessor, IPAttnProcessor, IPAttnProcessor2_0, \
|
||||
AttnProcessor2_0
|
||||
from ipadapter.ip_adapter.ip_adapter import ImageProjModel
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from toolkit.config_modules import AdapterConfig, AdapterTypes
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
import weakref
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPVisionModelWithProjection,
|
||||
CLIPVisionModel,
|
||||
AutoImageProcessor,
|
||||
ConvNextModel,
|
||||
ConvNextForImageClassification,
|
||||
ConvNextImageProcessor,
|
||||
UMT5EncoderModel, LlamaTokenizerFast
|
||||
)
|
||||
from toolkit.models.size_agnostic_feature_encoder import SAFEImageProcessor, SAFEVisionModel
|
||||
|
||||
from transformers import ViTHybridImageProcessor, ViTHybridForImageClassification
|
||||
|
||||
from transformers import ViTFeatureExtractor, ViTForImageClassification
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class CustomAdapter(torch.nn.Module):
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: 'AdapterConfig'):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.image_processor: CLIPImageProcessor = None
|
||||
self.input_size = 224
|
||||
self.adapter_type: AdapterTypes = self.config.type
|
||||
self.current_scale = 1.0
|
||||
self.is_active = True
|
||||
self.flag_word = "fla9wor0"
|
||||
self.is_unconditional_run = False
|
||||
|
||||
self.vision_encoder: Union[PhotoMakerCLIPEncoder, CLIPVisionModelWithProjection] = None
|
||||
|
||||
self.fuse_module: FuseModule = None
|
||||
|
||||
self.lora: None = None
|
||||
|
||||
self.position_ids: Optional[List[int]] = None
|
||||
|
||||
self.num_control_images = 1
|
||||
self.token_mask: Optional[torch.Tensor] = None
|
||||
|
||||
# setup clip
|
||||
self.setup_clip()
|
||||
# add for dataloader
|
||||
self.clip_image_processor = self.image_processor
|
||||
|
||||
self.clip_fusion_module: CLIPFusionModule = None
|
||||
self.ilora_module: InstantLoRAModule = None
|
||||
|
||||
self.te: Union[T5EncoderModel, CLIPTextModel] = None
|
||||
self.tokenizer: CLIPTokenizer = None
|
||||
self.te_adapter: TEAdapter = None
|
||||
self.te_augmenter: TEAugAdapter = None
|
||||
self.vd_adapter: VisionDirectAdapter = None
|
||||
self.single_value_adapter: SingleValueAdapter = None
|
||||
self.conditional_embeds: Optional[torch.Tensor] = None
|
||||
self.unconditional_embeds: Optional[torch.Tensor] = None
|
||||
|
||||
self.setup_adapter()
|
||||
|
||||
if self.adapter_type == 'photo_maker':
|
||||
# try to load from our name_or_path
|
||||
if self.config.name_or_path is not None and self.config.name_or_path.endswith('.bin'):
|
||||
self.load_state_dict(torch.load(self.config.name_or_path, map_location=self.device), strict=False)
|
||||
# add the trigger word to the tokenizer
|
||||
if isinstance(self.sd_ref().tokenizer, list):
|
||||
for tokenizer in self.sd_ref().tokenizer:
|
||||
tokenizer.add_tokens([self.flag_word], special_tokens=True)
|
||||
else:
|
||||
self.sd_ref().tokenizer.add_tokens([self.flag_word], special_tokens=True)
|
||||
elif self.config.name_or_path is not None:
|
||||
loaded_state_dict = load_custom_adapter_model(
|
||||
self.config.name_or_path,
|
||||
self.sd_ref().device,
|
||||
dtype=self.sd_ref().dtype,
|
||||
)
|
||||
self.load_state_dict(loaded_state_dict, strict=False)
|
||||
|
||||
def setup_adapter(self):
|
||||
torch_dtype = get_torch_dtype(self.sd_ref().dtype)
|
||||
if self.adapter_type == 'photo_maker':
|
||||
sd = self.sd_ref()
|
||||
embed_dim = sd.unet.config['cross_attention_dim']
|
||||
self.fuse_module = FuseModule(embed_dim)
|
||||
elif self.adapter_type == 'clip_fusion':
|
||||
sd = self.sd_ref()
|
||||
embed_dim = sd.unet.config['cross_attention_dim']
|
||||
|
||||
vision_tokens = ((self.vision_encoder.config.image_size // self.vision_encoder.config.patch_size) ** 2)
|
||||
if self.config.image_encoder_arch == 'clip':
|
||||
vision_tokens = vision_tokens + 1
|
||||
self.clip_fusion_module = CLIPFusionModule(
|
||||
text_hidden_size=embed_dim,
|
||||
text_tokens=77,
|
||||
vision_hidden_size=self.vision_encoder.config.hidden_size,
|
||||
vision_tokens=vision_tokens
|
||||
)
|
||||
elif self.adapter_type == 'ilora':
|
||||
vision_tokens = ((self.vision_encoder.config.image_size // self.vision_encoder.config.patch_size) ** 2)
|
||||
if self.config.image_encoder_arch == 'clip':
|
||||
vision_tokens = vision_tokens + 1
|
||||
|
||||
vision_hidden_size = self.vision_encoder.config.hidden_size
|
||||
|
||||
if self.config.clip_layer == 'image_embeds':
|
||||
vision_tokens = 1
|
||||
vision_hidden_size = self.vision_encoder.config.projection_dim
|
||||
|
||||
self.ilora_module = InstantLoRAModule(
|
||||
vision_tokens=vision_tokens,
|
||||
vision_hidden_size=vision_hidden_size,
|
||||
head_dim=self.config.head_dim,
|
||||
num_heads=self.config.num_heads,
|
||||
sd=self.sd_ref(),
|
||||
config=self.config
|
||||
)
|
||||
elif self.adapter_type == 'text_encoder':
|
||||
if self.config.text_encoder_arch == 't5':
|
||||
te_kwargs = {}
|
||||
# te_kwargs['load_in_4bit'] = True
|
||||
# te_kwargs['load_in_8bit'] = True
|
||||
te_kwargs['device_map'] = "auto"
|
||||
te_is_quantized = True
|
||||
|
||||
self.te = T5EncoderModel.from_pretrained(
|
||||
self.config.text_encoder_path,
|
||||
torch_dtype=torch_dtype,
|
||||
**te_kwargs
|
||||
)
|
||||
|
||||
# self.te.to = lambda *args, **kwargs: None
|
||||
self.tokenizer = T5Tokenizer.from_pretrained(self.config.text_encoder_path)
|
||||
elif self.config.text_encoder_arch == 'pile-t5':
|
||||
te_kwargs = {}
|
||||
# te_kwargs['load_in_4bit'] = True
|
||||
# te_kwargs['load_in_8bit'] = True
|
||||
te_kwargs['device_map'] = "auto"
|
||||
te_is_quantized = True
|
||||
|
||||
self.te = UMT5EncoderModel.from_pretrained(
|
||||
self.config.text_encoder_path,
|
||||
torch_dtype=torch_dtype,
|
||||
**te_kwargs
|
||||
)
|
||||
|
||||
# self.te.to = lambda *args, **kwargs: None
|
||||
self.tokenizer = LlamaTokenizerFast.from_pretrained(self.config.text_encoder_path)
|
||||
if self.tokenizer.pad_token is None:
|
||||
self.tokenizer.add_special_tokens({'pad_token': '[PAD]'})
|
||||
elif self.config.text_encoder_arch == 'clip':
|
||||
self.te = CLIPTextModel.from_pretrained(self.config.text_encoder_path).to(self.sd_ref().unet.device,
|
||||
dtype=torch_dtype)
|
||||
self.tokenizer = CLIPTokenizer.from_pretrained(self.config.text_encoder_path)
|
||||
else:
|
||||
raise ValueError(f"unknown text encoder arch: {self.config.text_encoder_arch}")
|
||||
|
||||
self.te_adapter = TEAdapter(self, self.sd_ref(), self.te, self.tokenizer)
|
||||
elif self.adapter_type == 'te_augmenter':
|
||||
self.te_augmenter = TEAugAdapter(self, self.sd_ref())
|
||||
elif self.adapter_type == 'vision_direct':
|
||||
self.vd_adapter = VisionDirectAdapter(self, self.sd_ref(), self.vision_encoder)
|
||||
elif self.adapter_type == 'single_value':
|
||||
self.single_value_adapter = SingleValueAdapter(self, self.sd_ref(), num_values=self.config.num_tokens)
|
||||
else:
|
||||
raise ValueError(f"unknown adapter type: {self.adapter_type}")
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
# dont think this is used
|
||||
# if self.adapter_type == 'photo_maker':
|
||||
# id_pixel_values = args[0]
|
||||
# prompt_embeds: PromptEmbeds = args[1]
|
||||
# class_tokens_mask = args[2]
|
||||
#
|
||||
# grads_on_image_encoder = self.config.train_image_encoder and torch.is_grad_enabled()
|
||||
#
|
||||
# with torch.set_grad_enabled(grads_on_image_encoder):
|
||||
# id_embeds = self.vision_encoder(self, id_pixel_values, do_projection2=False)
|
||||
#
|
||||
# if not grads_on_image_encoder:
|
||||
# id_embeds = id_embeds.detach()
|
||||
#
|
||||
# prompt_embeds = prompt_embeds.detach()
|
||||
#
|
||||
# updated_prompt_embeds = self.fuse_module(
|
||||
# prompt_embeds, id_embeds, class_tokens_mask
|
||||
# )
|
||||
#
|
||||
# return updated_prompt_embeds
|
||||
# else:
|
||||
raise NotImplementedError
|
||||
|
||||
def setup_clip(self):
|
||||
adapter_config = self.config
|
||||
sd = self.sd_ref()
|
||||
if self.config.type == "text_encoder" or self.config.type == "single_value":
|
||||
return
|
||||
if self.config.type == 'photo_maker':
|
||||
try:
|
||||
self.image_processor = CLIPImageProcessor.from_pretrained(self.config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = CLIPImageProcessor()
|
||||
if self.config.image_encoder_path is None:
|
||||
self.vision_encoder = PhotoMakerCLIPEncoder()
|
||||
else:
|
||||
self.vision_encoder = PhotoMakerCLIPEncoder.from_pretrained(self.config.image_encoder_path)
|
||||
elif self.config.image_encoder_arch == 'clip' or self.config.image_encoder_arch == 'clip+':
|
||||
try:
|
||||
self.image_processor = CLIPImageProcessor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = CLIPImageProcessor()
|
||||
self.vision_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
adapter_config.image_encoder_path,
|
||||
ignore_mismatched_sizes=True).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
elif self.config.image_encoder_arch == 'siglip':
|
||||
from transformers import SiglipImageProcessor, SiglipVisionModel
|
||||
try:
|
||||
self.image_processor = SiglipImageProcessor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = SiglipImageProcessor()
|
||||
self.vision_encoder = SiglipVisionModel.from_pretrained(
|
||||
adapter_config.image_encoder_path,
|
||||
ignore_mismatched_sizes=True).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
elif self.config.image_encoder_arch == 'vit':
|
||||
try:
|
||||
self.image_processor = ViTFeatureExtractor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = ViTFeatureExtractor()
|
||||
self.vision_encoder = ViTForImageClassification.from_pretrained(adapter_config.image_encoder_path).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
elif self.config.image_encoder_arch == 'safe':
|
||||
try:
|
||||
self.image_processor = SAFEImageProcessor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = SAFEImageProcessor()
|
||||
self.vision_encoder = SAFEVisionModel(
|
||||
in_channels=3,
|
||||
num_tokens=self.config.safe_tokens,
|
||||
num_vectors=sd.unet.config['cross_attention_dim'],
|
||||
reducer_channels=self.config.safe_reducer_channels,
|
||||
channels=self.config.safe_channels,
|
||||
downscale_factor=8
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
elif self.config.image_encoder_arch == 'convnext':
|
||||
try:
|
||||
self.image_processor = ConvNextImageProcessor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
print(f"could not load image processor from {adapter_config.image_encoder_path}")
|
||||
self.image_processor = ConvNextImageProcessor(
|
||||
size=320,
|
||||
image_mean=[0.48145466, 0.4578275, 0.40821073],
|
||||
image_std=[0.26862954, 0.26130258, 0.27577711],
|
||||
)
|
||||
self.vision_encoder = ConvNextForImageClassification.from_pretrained(
|
||||
adapter_config.image_encoder_path,
|
||||
use_safetensors=True,
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
elif self.config.image_encoder_arch == 'vit-hybrid':
|
||||
try:
|
||||
self.image_processor = ViTHybridImageProcessor.from_pretrained(adapter_config.image_encoder_path)
|
||||
except EnvironmentError:
|
||||
print(f"could not load image processor from {adapter_config.image_encoder_path}")
|
||||
self.image_processor = ViTHybridImageProcessor(
|
||||
size=320,
|
||||
image_mean=[0.48145466, 0.4578275, 0.40821073],
|
||||
image_std=[0.26862954, 0.26130258, 0.27577711],
|
||||
)
|
||||
self.vision_encoder = ViTHybridForImageClassification.from_pretrained(
|
||||
adapter_config.image_encoder_path,
|
||||
use_safetensors=True,
|
||||
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
else:
|
||||
raise ValueError(f"unknown image encoder arch: {adapter_config.image_encoder_arch}")
|
||||
|
||||
self.input_size = self.vision_encoder.config.image_size
|
||||
|
||||
if self.config.quad_image: # 4x4 image
|
||||
# self.clip_image_processor.config
|
||||
# We do a 3x downscale of the image, so we need to adjust the input size
|
||||
preprocessor_input_size = self.vision_encoder.config.image_size * 2
|
||||
|
||||
# update the preprocessor so images come in at the right size
|
||||
if 'height' in self.image_processor.size:
|
||||
self.image_processor.size['height'] = preprocessor_input_size
|
||||
self.image_processor.size['width'] = preprocessor_input_size
|
||||
elif hasattr(self.image_processor, 'crop_size'):
|
||||
self.image_processor.size['shortest_edge'] = preprocessor_input_size
|
||||
self.image_processor.crop_size['height'] = preprocessor_input_size
|
||||
self.image_processor.crop_size['width'] = preprocessor_input_size
|
||||
|
||||
if self.config.image_encoder_arch == 'clip+':
|
||||
# self.image_processor.config
|
||||
# We do a 3x downscale of the image, so we need to adjust the input size
|
||||
preprocessor_input_size = self.vision_encoder.config.image_size * 4
|
||||
|
||||
# update the preprocessor so images come in at the right size
|
||||
self.image_processor.size['shortest_edge'] = preprocessor_input_size
|
||||
self.image_processor.crop_size['height'] = preprocessor_input_size
|
||||
self.image_processor.crop_size['width'] = preprocessor_input_size
|
||||
|
||||
self.preprocessor = CLIPImagePreProcessor(
|
||||
input_size=preprocessor_input_size,
|
||||
clip_input_size=self.vision_encoder.config.image_size,
|
||||
)
|
||||
if 'height' in self.image_processor.size:
|
||||
self.input_size = self.image_processor.size['height']
|
||||
else:
|
||||
self.input_size = self.image_processor.crop_size['height']
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||
strict = False
|
||||
if self.config.train_only_image_encoder and 'vd_adapter' not in state_dict and 'dvadapter' not in state_dict:
|
||||
# we are loading pure clip weights.
|
||||
self.vision_encoder.load_state_dict(state_dict, strict=strict)
|
||||
|
||||
if 'lora_weights' in state_dict:
|
||||
# todo add LoRA
|
||||
# self.sd_ref().pipeline.load_lora_weights(state_dict["lora_weights"], adapter_name="photomaker")
|
||||
# self.sd_ref().pipeline.fuse_lora()
|
||||
pass
|
||||
if 'clip_fusion' in state_dict:
|
||||
self.clip_fusion_module.load_state_dict(state_dict['clip_fusion'], strict=strict)
|
||||
if 'id_encoder' in state_dict and (self.adapter_type == 'photo_maker' or self.adapter_type == 'clip_fusion'):
|
||||
self.vision_encoder.load_state_dict(state_dict['id_encoder'], strict=strict)
|
||||
# check to see if the fuse weights are there
|
||||
fuse_weights = {}
|
||||
for k, v in state_dict['id_encoder'].items():
|
||||
if k.startswith('fuse_module'):
|
||||
k = k.replace('fuse_module.', '')
|
||||
fuse_weights[k] = v
|
||||
if len(fuse_weights) > 0:
|
||||
try:
|
||||
self.fuse_module.load_state_dict(fuse_weights, strict=strict)
|
||||
except Exception as e:
|
||||
|
||||
print(e)
|
||||
# force load it
|
||||
print(f"force loading fuse module as it did not match")
|
||||
current_state_dict = self.fuse_module.state_dict()
|
||||
for k, v in fuse_weights.items():
|
||||
if len(v.shape) == 1:
|
||||
current_state_dict[k] = v[:current_state_dict[k].shape[0]]
|
||||
elif len(v.shape) == 2:
|
||||
current_state_dict[k] = v[:current_state_dict[k].shape[0], :current_state_dict[k].shape[1]]
|
||||
elif len(v.shape) == 3:
|
||||
current_state_dict[k] = v[:current_state_dict[k].shape[0], :current_state_dict[k].shape[1],
|
||||
:current_state_dict[k].shape[2]]
|
||||
elif len(v.shape) == 4:
|
||||
current_state_dict[k] = v[:current_state_dict[k].shape[0], :current_state_dict[k].shape[1],
|
||||
:current_state_dict[k].shape[2], :current_state_dict[k].shape[3]]
|
||||
else:
|
||||
raise ValueError(f"unknown shape: {v.shape}")
|
||||
self.fuse_module.load_state_dict(current_state_dict, strict=strict)
|
||||
|
||||
if 'te_adapter' in state_dict:
|
||||
self.te_adapter.load_state_dict(state_dict['te_adapter'], strict=strict)
|
||||
|
||||
if 'te_augmenter' in state_dict:
|
||||
self.te_augmenter.load_state_dict(state_dict['te_augmenter'], strict=strict)
|
||||
|
||||
if 'vd_adapter' in state_dict:
|
||||
self.vd_adapter.load_state_dict(state_dict['vd_adapter'], strict=strict)
|
||||
if 'dvadapter' in state_dict:
|
||||
self.vd_adapter.load_state_dict(state_dict['dvadapter'], strict=strict)
|
||||
|
||||
if 'sv_adapter' in state_dict:
|
||||
self.single_value_adapter.load_state_dict(state_dict['sv_adapter'], strict=strict)
|
||||
|
||||
if 'vision_encoder' in state_dict and self.config.train_image_encoder:
|
||||
self.vision_encoder.load_state_dict(state_dict['vision_encoder'], strict=strict)
|
||||
|
||||
if 'fuse_module' in state_dict:
|
||||
self.fuse_module.load_state_dict(state_dict['fuse_module'], strict=strict)
|
||||
|
||||
if 'ilora' in state_dict:
|
||||
try:
|
||||
self.ilora_module.load_state_dict(state_dict['ilora'], strict=strict)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
pass
|
||||
|
||||
def state_dict(self) -> OrderedDict:
|
||||
state_dict = OrderedDict()
|
||||
if self.config.train_only_image_encoder:
|
||||
return self.vision_encoder.state_dict()
|
||||
|
||||
if self.adapter_type == 'photo_maker':
|
||||
if self.config.train_image_encoder:
|
||||
state_dict["id_encoder"] = self.vision_encoder.state_dict()
|
||||
|
||||
state_dict["fuse_module"] = self.fuse_module.state_dict()
|
||||
|
||||
# todo save LoRA
|
||||
return state_dict
|
||||
|
||||
elif self.adapter_type == 'clip_fusion':
|
||||
if self.config.train_image_encoder:
|
||||
state_dict["vision_encoder"] = self.vision_encoder.state_dict()
|
||||
state_dict["clip_fusion"] = self.clip_fusion_module.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'text_encoder':
|
||||
state_dict["te_adapter"] = self.te_adapter.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'te_augmenter':
|
||||
if self.config.train_image_encoder:
|
||||
state_dict["vision_encoder"] = self.vision_encoder.state_dict()
|
||||
state_dict["te_augmenter"] = self.te_augmenter.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'vision_direct':
|
||||
state_dict["dvadapter"] = self.vd_adapter.state_dict()
|
||||
if self.config.train_image_encoder:
|
||||
state_dict["vision_encoder"] = self.vision_encoder.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'single_value':
|
||||
state_dict["sv_adapter"] = self.single_value_adapter.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'ilora':
|
||||
if self.config.train_image_encoder:
|
||||
state_dict["vision_encoder"] = self.vision_encoder.state_dict()
|
||||
state_dict["ilora"] = self.ilora_module.state_dict()
|
||||
return state_dict
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def add_extra_values(self, extra_values: torch.Tensor, is_unconditional=False):
|
||||
if self.adapter_type == 'single_value':
|
||||
if is_unconditional:
|
||||
self.unconditional_embeds = extra_values.to(self.device, get_torch_dtype(self.sd_ref().dtype))
|
||||
else:
|
||||
self.conditional_embeds = extra_values.to(self.device, get_torch_dtype(self.sd_ref().dtype))
|
||||
|
||||
|
||||
def condition_prompt(
|
||||
self,
|
||||
prompt: Union[List[str], str],
|
||||
is_unconditional: bool = False,
|
||||
):
|
||||
if self.adapter_type == 'clip_fusion' or self.adapter_type == 'ilora' or self.adapter_type == 'vision_direct':
|
||||
return prompt
|
||||
elif self.adapter_type == 'text_encoder':
|
||||
# todo allow for training
|
||||
with torch.no_grad():
|
||||
# encode and save the embeds
|
||||
if is_unconditional:
|
||||
self.unconditional_embeds = self.te_adapter.encode_text(prompt).detach()
|
||||
else:
|
||||
self.conditional_embeds = self.te_adapter.encode_text(prompt).detach()
|
||||
return prompt
|
||||
elif self.adapter_type == 'photo_maker':
|
||||
if is_unconditional:
|
||||
return prompt
|
||||
else:
|
||||
|
||||
with torch.no_grad():
|
||||
was_list = isinstance(prompt, list)
|
||||
if not was_list:
|
||||
prompt_list = [prompt]
|
||||
else:
|
||||
prompt_list = prompt
|
||||
|
||||
new_prompt_list = []
|
||||
token_mask_list = []
|
||||
|
||||
for prompt in prompt_list:
|
||||
|
||||
our_class = None
|
||||
# find a class in the prompt
|
||||
prompt_parts = prompt.split(' ')
|
||||
prompt_parts = [p.strip().lower() for p in prompt_parts if len(p) > 0]
|
||||
|
||||
new_prompt_parts = []
|
||||
tokened_prompt_parts = []
|
||||
for idx, prompt_part in enumerate(prompt_parts):
|
||||
new_prompt_parts.append(prompt_part)
|
||||
tokened_prompt_parts.append(prompt_part)
|
||||
if prompt_part in self.config.class_names:
|
||||
our_class = prompt_part
|
||||
# add the flag word
|
||||
tokened_prompt_parts.append(self.flag_word)
|
||||
|
||||
if self.num_control_images > 1:
|
||||
# add the rest
|
||||
for _ in range(self.num_control_images - 1):
|
||||
new_prompt_parts.extend(prompt_parts[idx + 1:])
|
||||
|
||||
# add the rest
|
||||
tokened_prompt_parts.extend(prompt_parts[idx + 1:])
|
||||
new_prompt_parts.extend(prompt_parts[idx + 1:])
|
||||
|
||||
break
|
||||
|
||||
prompt = " ".join(new_prompt_parts)
|
||||
tokened_prompt = " ".join(tokened_prompt_parts)
|
||||
|
||||
if our_class is None:
|
||||
# add the first one to the front of the prompt
|
||||
tokened_prompt = self.config.class_names[0] + ' ' + self.flag_word + ' ' + prompt
|
||||
our_class = self.config.class_names[0]
|
||||
prompt = " ".join(
|
||||
[self.config.class_names[0] for _ in range(self.num_control_images)]) + ' ' + prompt
|
||||
|
||||
# add the prompt to the list
|
||||
new_prompt_list.append(prompt)
|
||||
|
||||
# tokenize them with just the first tokenizer
|
||||
tokenizer = self.sd_ref().tokenizer
|
||||
if isinstance(tokenizer, list):
|
||||
tokenizer = tokenizer[0]
|
||||
|
||||
flag_token = tokenizer.convert_tokens_to_ids(self.flag_word)
|
||||
|
||||
tokenized_prompt = tokenizer.encode(prompt)
|
||||
tokenized_tokened_prompt = tokenizer.encode(tokened_prompt)
|
||||
|
||||
flag_idx = tokenized_tokened_prompt.index(flag_token)
|
||||
|
||||
class_token = tokenized_prompt[flag_idx - 1]
|
||||
|
||||
boolean_mask = torch.zeros(flag_idx - 1, dtype=torch.bool)
|
||||
boolean_mask = torch.cat((boolean_mask, torch.ones(self.num_control_images, dtype=torch.bool)))
|
||||
boolean_mask = boolean_mask.to(self.device)
|
||||
# zero pad it to 77
|
||||
boolean_mask = F.pad(boolean_mask, (0, 77 - boolean_mask.shape[0]), value=False)
|
||||
|
||||
token_mask_list.append(boolean_mask)
|
||||
|
||||
self.token_mask = torch.cat(token_mask_list, dim=0).to(self.device)
|
||||
|
||||
prompt_list = new_prompt_list
|
||||
|
||||
if not was_list:
|
||||
prompt = prompt_list[0]
|
||||
else:
|
||||
prompt = prompt_list
|
||||
|
||||
return prompt
|
||||
|
||||
else:
|
||||
return prompt
|
||||
|
||||
def condition_encoded_embeds(
|
||||
self,
|
||||
tensors_0_1: torch.Tensor,
|
||||
prompt_embeds: PromptEmbeds,
|
||||
is_training=False,
|
||||
has_been_preprocessed=False,
|
||||
is_unconditional=False,
|
||||
quad_count=4,
|
||||
is_generating_samples=False,
|
||||
) -> PromptEmbeds:
|
||||
if self.adapter_type == 'text_encoder' and is_generating_samples:
|
||||
# replace the prompt embed with ours
|
||||
if is_unconditional:
|
||||
return self.unconditional_embeds.clone()
|
||||
return self.conditional_embeds.clone()
|
||||
|
||||
if self.adapter_type == 'ilora':
|
||||
return prompt_embeds
|
||||
|
||||
if self.adapter_type == 'photo_maker' or self.adapter_type == 'clip_fusion':
|
||||
if is_unconditional:
|
||||
# we dont condition the negative embeds for photo maker
|
||||
return prompt_embeds.clone()
|
||||
with torch.no_grad():
|
||||
# on training the clip image is created in the dataloader
|
||||
if not has_been_preprocessed:
|
||||
# tensors should be 0-1
|
||||
if tensors_0_1.ndim == 3:
|
||||
tensors_0_1 = tensors_0_1.unsqueeze(0)
|
||||
# training tensors are 0 - 1
|
||||
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
|
||||
# if images are out of this range throw error
|
||||
if tensors_0_1.min() < -0.3 or tensors_0_1.max() > 1.3:
|
||||
raise ValueError("image tensor values must be between 0 and 1. Got min: {}, max: {}".format(
|
||||
tensors_0_1.min(), tensors_0_1.max()
|
||||
))
|
||||
clip_image = self.image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
clip_image = clip_image.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype)).detach()
|
||||
|
||||
if self.config.quad_image:
|
||||
# split the 4x4 grid and stack on batch
|
||||
ci1, ci2 = clip_image.chunk(2, dim=2)
|
||||
ci1, ci3 = ci1.chunk(2, dim=3)
|
||||
ci2, ci4 = ci2.chunk(2, dim=3)
|
||||
to_cat = []
|
||||
for i, ci in enumerate([ci1, ci2, ci3, ci4]):
|
||||
if i < quad_count:
|
||||
to_cat.append(ci)
|
||||
else:
|
||||
break
|
||||
|
||||
clip_image = torch.cat(to_cat, dim=0).detach()
|
||||
|
||||
if self.adapter_type == 'photo_maker':
|
||||
# Embeddings need to be (b, num_inputs, c, h, w) for now, just put 1 input image
|
||||
clip_image = clip_image.unsqueeze(1)
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training and self.config.train_image_encoder:
|
||||
self.vision_encoder.train()
|
||||
clip_image = clip_image.requires_grad_(True)
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image,
|
||||
do_projection2=isinstance(self.sd_ref().text_encoder, list),
|
||||
)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
self.vision_encoder.eval()
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image, do_projection2=isinstance(self.sd_ref().text_encoder, list)
|
||||
).detach()
|
||||
|
||||
prompt_embeds.text_embeds = self.fuse_module(
|
||||
prompt_embeds.text_embeds,
|
||||
id_embeds,
|
||||
self.token_mask
|
||||
)
|
||||
return prompt_embeds
|
||||
elif self.adapter_type == 'clip_fusion':
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training and self.config.train_image_encoder:
|
||||
self.vision_encoder.train()
|
||||
clip_image = clip_image.requires_grad_(True)
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
self.vision_encoder.eval()
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image, output_hidden_states=True
|
||||
)
|
||||
|
||||
img_embeds = id_embeds['last_hidden_state']
|
||||
|
||||
if self.config.quad_image:
|
||||
# get the outputs of the quat
|
||||
chunks = img_embeds.chunk(quad_count, dim=0)
|
||||
chunk_sum = torch.zeros_like(chunks[0])
|
||||
for chunk in chunks:
|
||||
chunk_sum = chunk_sum + chunk
|
||||
# get the mean of them
|
||||
|
||||
img_embeds = chunk_sum / quad_count
|
||||
|
||||
if not is_training or not self.config.train_image_encoder:
|
||||
img_embeds = img_embeds.detach()
|
||||
|
||||
prompt_embeds.text_embeds = self.clip_fusion_module(
|
||||
prompt_embeds.text_embeds,
|
||||
img_embeds
|
||||
)
|
||||
return prompt_embeds
|
||||
|
||||
|
||||
else:
|
||||
return prompt_embeds
|
||||
|
||||
def get_empty_clip_image(self, batch_size: int) -> torch.Tensor:
|
||||
with torch.no_grad():
|
||||
tensors_0_1 = torch.rand([batch_size, 3, self.input_size, self.input_size], device=self.device)
|
||||
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
|
||||
dtype=get_torch_dtype(self.sd_ref().dtype))
|
||||
tensors_0_1 = tensors_0_1 * noise_scale
|
||||
# tensors_0_1 = tensors_0_1 * 0
|
||||
mean = torch.tensor(self.clip_image_processor.image_mean).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
std = torch.tensor(self.clip_image_processor.image_std).to(
|
||||
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
|
||||
).detach()
|
||||
tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
|
||||
clip_image = (tensors_0_1 - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
|
||||
return clip_image.detach()
|
||||
|
||||
def train(self, mode: bool = True):
|
||||
if self.config.train_image_encoder:
|
||||
self.vision_encoder.train(mode)
|
||||
else:
|
||||
super().train(mode)
|
||||
|
||||
def trigger_pre_te(
|
||||
self,
|
||||
tensors_0_1: torch.Tensor,
|
||||
is_training=False,
|
||||
has_been_preprocessed=False,
|
||||
quad_count=4,
|
||||
batch_size=1,
|
||||
) -> PromptEmbeds:
|
||||
if self.adapter_type == 'ilora' or self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter':
|
||||
if tensors_0_1 is None:
|
||||
tensors_0_1 = self.get_empty_clip_image(batch_size)
|
||||
has_been_preprocessed = True
|
||||
|
||||
with torch.no_grad():
|
||||
# on training the clip image is created in the dataloader
|
||||
if not has_been_preprocessed:
|
||||
# tensors should be 0-1
|
||||
if tensors_0_1.ndim == 3:
|
||||
tensors_0_1 = tensors_0_1.unsqueeze(0)
|
||||
# training tensors are 0 - 1
|
||||
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
|
||||
# if images are out of this range throw error
|
||||
if tensors_0_1.min() < -0.3 or tensors_0_1.max() > 1.3:
|
||||
raise ValueError("image tensor values must be between 0 and 1. Got min: {}, max: {}".format(
|
||||
tensors_0_1.min(), tensors_0_1.max()
|
||||
))
|
||||
clip_image = self.image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
|
||||
batch_size = clip_image.shape[0]
|
||||
if self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter':
|
||||
# add an unconditional so we can save it
|
||||
unconditional = self.get_empty_clip_image(batch_size).to(
|
||||
clip_image.device, dtype=clip_image.dtype
|
||||
)
|
||||
clip_image = torch.cat([unconditional, clip_image], dim=0)
|
||||
|
||||
clip_image = clip_image.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype)).detach()
|
||||
|
||||
if self.config.quad_image:
|
||||
# split the 4x4 grid and stack on batch
|
||||
ci1, ci2 = clip_image.chunk(2, dim=2)
|
||||
ci1, ci3 = ci1.chunk(2, dim=3)
|
||||
ci2, ci4 = ci2.chunk(2, dim=3)
|
||||
to_cat = []
|
||||
for i, ci in enumerate([ci1, ci2, ci3, ci4]):
|
||||
if i < quad_count:
|
||||
to_cat.append(ci)
|
||||
else:
|
||||
break
|
||||
|
||||
clip_image = torch.cat(to_cat, dim=0).detach()
|
||||
|
||||
if self.adapter_type == 'ilora':
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training and self.config.train_image_encoder:
|
||||
self.vision_encoder.train()
|
||||
clip_image = clip_image.requires_grad_(True)
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
self.vision_encoder.eval()
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image, output_hidden_states=True
|
||||
)
|
||||
|
||||
if self.config.clip_layer == 'penultimate_hidden_states':
|
||||
img_embeds = id_embeds.hidden_states[-2]
|
||||
elif self.config.clip_layer == 'last_hidden_state':
|
||||
img_embeds = id_embeds.hidden_states[-1]
|
||||
elif self.config.clip_layer == 'image_embeds':
|
||||
img_embeds = id_embeds.image_embeds
|
||||
else:
|
||||
raise ValueError(f"unknown clip layer: {self.config.clip_layer}")
|
||||
|
||||
if self.config.quad_image:
|
||||
# get the outputs of the quat
|
||||
chunks = img_embeds.chunk(quad_count, dim=0)
|
||||
chunk_sum = torch.zeros_like(chunks[0])
|
||||
for chunk in chunks:
|
||||
chunk_sum = chunk_sum + chunk
|
||||
# get the mean of them
|
||||
|
||||
img_embeds = chunk_sum / quad_count
|
||||
|
||||
if not is_training or not self.config.train_image_encoder:
|
||||
img_embeds = img_embeds.detach()
|
||||
|
||||
self.ilora_module(img_embeds)
|
||||
if self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter':
|
||||
with torch.set_grad_enabled(is_training):
|
||||
if is_training and self.config.train_image_encoder:
|
||||
self.vision_encoder.train()
|
||||
clip_image = clip_image.requires_grad_(True)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
self.vision_encoder.eval()
|
||||
clip_output = self.vision_encoder(
|
||||
clip_image,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
if self.config.clip_layer == 'penultimate_hidden_states':
|
||||
# they skip last layer for ip+
|
||||
# https://github.com/tencent-ailab/IP-Adapter/blob/f4b6742db35ea6d81c7b829a55b0a312c7f5a677/tutorial_train_plus.py#L403C26-L403C26
|
||||
clip_image_embeds = clip_output.hidden_states[-2]
|
||||
elif self.config.clip_layer == 'last_hidden_state':
|
||||
clip_image_embeds = clip_output.hidden_states[-1]
|
||||
else:
|
||||
clip_image_embeds = clip_output.image_embeds
|
||||
# TODO should we always norm image embeds?
|
||||
# get norm embeddings
|
||||
l2_norm = torch.norm(clip_image_embeds, p=2)
|
||||
clip_image_embeds = clip_image_embeds / l2_norm
|
||||
|
||||
if not is_training or not self.config.train_image_encoder:
|
||||
clip_image_embeds = clip_image_embeds.detach()
|
||||
|
||||
if self.adapter_type == 'te_augmenter':
|
||||
clip_image_embeds = self.te_augmenter(clip_image_embeds)
|
||||
|
||||
if self.adapter_type == 'vision_direct':
|
||||
clip_image_embeds = self.vd_adapter(clip_image_embeds)
|
||||
|
||||
# save them to the conditional and unconditional
|
||||
try:
|
||||
self.unconditional_embeds, self.conditional_embeds = clip_image_embeds.chunk(2, dim=0)
|
||||
except ValueError:
|
||||
raise ValueError(f"could not split the clip image embeds into 2. Got shape: {clip_image_embeds.shape}")
|
||||
|
||||
def parameters(self, recurse: bool = True) -> Iterator[Parameter]:
|
||||
if self.config.train_only_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
return
|
||||
if self.config.type == 'photo_maker':
|
||||
yield from self.fuse_module.parameters(recurse)
|
||||
if self.config.train_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
elif self.config.type == 'clip_fusion':
|
||||
yield from self.clip_fusion_module.parameters(recurse)
|
||||
if self.config.train_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
elif self.config.type == 'ilora':
|
||||
yield from self.ilora_module.parameters(recurse)
|
||||
if self.config.train_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
elif self.config.type == 'text_encoder':
|
||||
for attn_processor in self.te_adapter.adapter_modules:
|
||||
yield from attn_processor.parameters(recurse)
|
||||
elif self.config.type == 'vision_direct':
|
||||
for attn_processor in self.vd_adapter.adapter_modules:
|
||||
yield from attn_processor.parameters(recurse)
|
||||
if self.config.train_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
elif self.config.type == 'te_augmenter':
|
||||
yield from self.te_augmenter.parameters(recurse)
|
||||
if self.config.train_image_encoder:
|
||||
yield from self.vision_encoder.parameters(recurse)
|
||||
elif self.config.type == 'single_value':
|
||||
yield from self.single_value_adapter.parameters(recurse)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
if hasattr(self.vision_encoder, "enable_gradient_checkpointing"):
|
||||
self.vision_encoder.enable_gradient_checkpointing()
|
||||
elif hasattr(self.vision_encoder, 'gradient_checkpointing'):
|
||||
self.vision_encoder.gradient_checkpointing = True
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
additional = {}
|
||||
if self.config.type == 'ilora':
|
||||
extra = self.ilora_module.get_additional_save_metadata()
|
||||
for k, v in extra.items():
|
||||
additional[k] = v
|
||||
additional['clip_layer'] = self.config.clip_layer
|
||||
additional['image_encoder_arch'] = self.config.head_dim
|
||||
return additional
|
||||
@@ -8,6 +8,7 @@ from typing import List, TYPE_CHECKING
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from PIL.ImageOps import exif_transpose
|
||||
from torchvision import transforms
|
||||
@@ -17,13 +18,57 @@ import albumentations as A
|
||||
|
||||
from toolkit.buckets import get_bucket_for_image_size, BucketResolution
|
||||
from toolkit.config_modules import DatasetConfig, preprocess_dataset_raw_config
|
||||
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments
|
||||
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments, CLIPCachingMixin
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
|
||||
import platform
|
||||
|
||||
def is_native_windows():
|
||||
return platform.system() == "Windows" and platform.release() != "2"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class RescaleTransform:
|
||||
"""Transform to rescale images to the range [-1, 1]."""
|
||||
|
||||
def __call__(self, image):
|
||||
return image * 2 - 1
|
||||
|
||||
|
||||
class NormalizeSDXLTransform:
|
||||
"""
|
||||
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
|
||||
|
||||
Mean: tensor([ 0.0002, -0.1034, -0.1879])
|
||||
Standard Deviation: tensor([0.5436, 0.5116, 0.5033])
|
||||
"""
|
||||
|
||||
def __call__(self, image):
|
||||
return transforms.Normalize(
|
||||
mean=[0.0002, -0.1034, -0.1879],
|
||||
std=[0.5436, 0.5116, 0.5033],
|
||||
)(image)
|
||||
|
||||
|
||||
class NormalizeSD15Transform:
|
||||
"""
|
||||
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
|
||||
|
||||
Mean: tensor([-0.1600, -0.2450, -0.3227])
|
||||
Standard Deviation: tensor([0.5319, 0.4997, 0.5139])
|
||||
|
||||
"""
|
||||
|
||||
def __call__(self, image):
|
||||
return transforms.Normalize(
|
||||
mean=[-0.1600, -0.2450, -0.3227],
|
||||
std=[0.5319, 0.4997, 0.5139],
|
||||
)(image)
|
||||
|
||||
|
||||
|
||||
class ImageDataset(Dataset, CaptionMixin):
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
@@ -63,7 +108,7 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
def get_config(self, key, default=None, required=False):
|
||||
@@ -80,7 +125,13 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
|
||||
def __getitem__(self, index):
|
||||
img_path = self.file_list[index]
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
try:
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
except Exception as e:
|
||||
print(f"Error opening image: {img_path}")
|
||||
print(e)
|
||||
# make a noise image if we can't open it
|
||||
img = Image.fromarray(np.random.randint(0, 255, (1024, 1024, 3), dtype=np.uint8))
|
||||
|
||||
# Downscale the source image first
|
||||
img = img.resize((int(img.size[0] * self.scale), int(img.size[1] * self.scale)), Image.BICUBIC)
|
||||
@@ -200,7 +251,7 @@ class PairedImageDataset(Dataset):
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
def get_all_prompts(self):
|
||||
@@ -315,7 +366,7 @@ class PairedImageDataset(Dataset):
|
||||
return img, prompt, (self.neg_weight, self.pos_weight)
|
||||
|
||||
|
||||
class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -333,6 +384,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
self.is_caching_latents = dataset_config.cache_latents or dataset_config.cache_latents_to_disk
|
||||
self.is_caching_latents_to_memory = dataset_config.cache_latents
|
||||
self.is_caching_latents_to_disk = dataset_config.cache_latents_to_disk
|
||||
self.is_caching_clip_vision_to_disk = dataset_config.cache_clip_vision_to_disk
|
||||
self.epoch_num = 0
|
||||
|
||||
self.sd = sd
|
||||
@@ -353,10 +405,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
|
||||
# check if dataset_path is a folder or json
|
||||
if os.path.isdir(self.dataset_path):
|
||||
file_list = [
|
||||
os.path.join(self.dataset_path, file) for file in os.listdir(self.dataset_path) if
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))
|
||||
]
|
||||
file_list = [os.path.join(root, file) for root, _, files in os.walk(self.dataset_path) for file in files if file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
|
||||
else:
|
||||
# assume json
|
||||
with open(self.dataset_path, 'r') as f:
|
||||
@@ -368,14 +417,45 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
# repeat the list
|
||||
file_list = file_list * self.dataset_config.num_repeats
|
||||
|
||||
if self.dataset_config.standardize_images:
|
||||
if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd:
|
||||
NormalizeMethod = NormalizeSDXLTransform
|
||||
else:
|
||||
NormalizeMethod = NormalizeSD15Transform
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
RescaleTransform(),
|
||||
NormalizeMethod(),
|
||||
])
|
||||
else:
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
RescaleTransform(),
|
||||
])
|
||||
|
||||
# this might take a while
|
||||
print(f"Dataset: {self.dataset_path}")
|
||||
print(f" - Preprocessing image dimensions")
|
||||
dataset_folder = self.dataset_path
|
||||
if not os.path.isdir(self.dataset_path):
|
||||
dataset_folder = os.path.dirname(dataset_folder)
|
||||
dataset_size_file = os.path.join(dataset_folder, '.aitk_size.json')
|
||||
if os.path.exists(dataset_size_file):
|
||||
with open(dataset_size_file, 'r') as f:
|
||||
self.size_database = json.load(f)
|
||||
else:
|
||||
self.size_database = {}
|
||||
|
||||
bad_count = 0
|
||||
for file in tqdm(file_list):
|
||||
try:
|
||||
file_item = FileItemDTO(
|
||||
sd=self.sd,
|
||||
path=file,
|
||||
dataset_config=dataset_config
|
||||
dataset_config=dataset_config,
|
||||
dataloader_transforms=self.transform,
|
||||
size_database=self.size_database,
|
||||
)
|
||||
self.file_list.append(file_item)
|
||||
except Exception as e:
|
||||
@@ -384,6 +464,10 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
print(e)
|
||||
bad_count += 1
|
||||
|
||||
# save the size database
|
||||
with open(dataset_size_file, 'w') as f:
|
||||
json.dump(self.size_database, f)
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
# print(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
|
||||
@@ -411,10 +495,6 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
if self.dataset_config.flip_x or self.dataset_config.flip_y:
|
||||
print(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
|
||||
])
|
||||
|
||||
self.setup_epoch()
|
||||
|
||||
@@ -427,6 +507,8 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
|
||||
self.setup_buckets()
|
||||
if self.is_caching_latents:
|
||||
self.cache_latents_all_latents()
|
||||
if self.is_caching_clip_vision_to_disk:
|
||||
self.cache_clip_vision_to_disk()
|
||||
else:
|
||||
if self.dataset_config.poi is not None:
|
||||
# handle cropping to a specific point of interest
|
||||
@@ -509,6 +591,13 @@ def get_dataloader_from_datasets(
|
||||
|
||||
# check if is caching latents
|
||||
|
||||
dataloader_kwargs = {}
|
||||
|
||||
if is_native_windows():
|
||||
dataloader_kwargs['num_workers'] = 0
|
||||
else:
|
||||
dataloader_kwargs['num_workers'] = dataset_config_list[0].num_workers
|
||||
dataloader_kwargs['prefetch_factor'] = dataset_config_list[0].prefetch_factor
|
||||
|
||||
if has_buckets:
|
||||
# make sure they all have buckets
|
||||
@@ -521,15 +610,15 @@ def get_dataloader_from_datasets(
|
||||
drop_last=False,
|
||||
shuffle=True,
|
||||
collate_fn=dto_collation, # Use the custom collate function
|
||||
num_workers=4
|
||||
**dataloader_kwargs
|
||||
)
|
||||
else:
|
||||
data_loader = DataLoader(
|
||||
concatenated_dataset,
|
||||
batch_size=batch_size,
|
||||
shuffle=True,
|
||||
num_workers=4,
|
||||
collate_fn=dto_collation
|
||||
collate_fn=dto_collation,
|
||||
**dataloader_kwargs
|
||||
)
|
||||
return data_loader
|
||||
|
||||
@@ -556,3 +645,19 @@ def trigger_dataloader_setup_epoch(dataloader: DataLoader):
|
||||
if hasattr(sub_dataset, 'setup_epoch'):
|
||||
sub_dataset.setup_epoch()
|
||||
sub_dataset.len = None
|
||||
|
||||
def get_dataloader_datasets(dataloader: DataLoader):
|
||||
# hacky but needed because of different types of datasets and dataloaders
|
||||
if isinstance(dataloader.dataset, list):
|
||||
datasets = []
|
||||
for dataset in dataloader.dataset:
|
||||
if hasattr(dataset, 'datasets'):
|
||||
for sub_dataset in dataset.datasets:
|
||||
datasets.append(sub_dataset)
|
||||
else:
|
||||
datasets.append(dataset)
|
||||
return datasets
|
||||
elif hasattr(dataloader.dataset, 'datasets'):
|
||||
return dataloader.dataset.datasets
|
||||
else:
|
||||
return [dataloader.dataset]
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
import os
|
||||
import weakref
|
||||
from _weakref import ReferenceType
|
||||
from typing import TYPE_CHECKING, List, Union
|
||||
import torch
|
||||
import random
|
||||
@@ -8,10 +11,12 @@ from PIL.ImageOps import exif_transpose
|
||||
from toolkit import image_utils
|
||||
from toolkit.dataloader_mixins import CaptionProcessingDTOMixin, ImageProcessingDTOMixin, LatentCachingFileItemDTOMixin, \
|
||||
ControlFileItemDTOMixin, ArgBreakMixin, PoiFileItemDTOMixin, MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, \
|
||||
UnconditionalFileItemDTOMixin
|
||||
UnconditionalFileItemDTOMixin, ClipImageFileItemDTOMixin
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.config_modules import DatasetConfig
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
printed_messages = []
|
||||
|
||||
@@ -28,6 +33,7 @@ class FileItemDTO(
|
||||
CaptionProcessingDTOMixin,
|
||||
ImageProcessingDTOMixin,
|
||||
ControlFileItemDTOMixin,
|
||||
ClipImageFileItemDTOMixin,
|
||||
MaskFileItemDTOMixin,
|
||||
AugmentationFileItemDTOMixin,
|
||||
UnconditionalFileItemDTOMixin,
|
||||
@@ -35,18 +41,25 @@ class FileItemDTO(
|
||||
ArgBreakMixin,
|
||||
):
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.path = kwargs.get('path', None)
|
||||
self.path = kwargs.get('path', '')
|
||||
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
# process width and height
|
||||
try:
|
||||
w, h = image_utils.get_image_size(self.path)
|
||||
except image_utils.UnknownImageFormat:
|
||||
print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
|
||||
f'This process is faster for png, jpeg')
|
||||
img = exif_transpose(Image.open(self.path))
|
||||
h, w = img.size
|
||||
size_database = kwargs.get('size_database', {})
|
||||
filename = os.path.basename(self.path)
|
||||
if filename in size_database:
|
||||
w, h = size_database[filename]
|
||||
else:
|
||||
# process width and height
|
||||
try:
|
||||
w, h = image_utils.get_image_size(self.path)
|
||||
except image_utils.UnknownImageFormat:
|
||||
print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
|
||||
f'This process is faster for png, jpeg')
|
||||
img = exif_transpose(Image.open(self.path))
|
||||
h, w = img.size
|
||||
size_database[filename] = (w, h)
|
||||
self.width: int = w
|
||||
self.height: int = h
|
||||
self.dataloader_transforms = kwargs.get('dataloader_transforms', None)
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
# self.caption_path: str = kwargs.get('caption_path', None)
|
||||
@@ -62,6 +75,7 @@ class FileItemDTO(
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_x', False)
|
||||
self.augments: List[str] = self.dataset_config.augments
|
||||
self.loss_multiplier: float = self.dataset_config.loss_multiplier
|
||||
|
||||
self.network_weight: float = self.dataset_config.network_weight
|
||||
self.is_reg = self.dataset_config.is_reg
|
||||
@@ -71,6 +85,7 @@ class FileItemDTO(
|
||||
self.tensor = None
|
||||
self.cleanup_latent()
|
||||
self.cleanup_control()
|
||||
self.cleanup_clip_image()
|
||||
self.cleanup_mask()
|
||||
self.cleanup_unconditional()
|
||||
|
||||
@@ -83,11 +98,15 @@ class DataLoaderBatchDTO:
|
||||
self.tensor: Union[torch.Tensor, None] = None
|
||||
self.latents: Union[torch.Tensor, None] = None
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
self.clip_image_tensor: Union[torch.Tensor, None] = None
|
||||
self.mask_tensor: Union[torch.Tensor, None] = None
|
||||
self.unaugmented_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_latents: Union[torch.Tensor, None] = None
|
||||
self.clip_image_embeds: Union[List[dict], None] = None
|
||||
self.clip_image_embeds_unconditional: Union[List[dict], None] = None
|
||||
self.sigmas: Union[torch.Tensor, None] = None # can be added elseware and passed along training code
|
||||
self.extra_values: Union[torch.Tensor, None] = torch.tensor([x.extra_values for x in self.file_items]) if len(self.file_items[0].extra_values) > 0 else None
|
||||
if not is_latents_cached:
|
||||
# only return a tensor if latents are not cached
|
||||
self.tensor: torch.Tensor = torch.cat([x.tensor.unsqueeze(0) for x in self.file_items])
|
||||
@@ -113,6 +132,23 @@ class DataLoaderBatchDTO:
|
||||
control_tensors.append(x.control_tensor)
|
||||
self.control_tensor = torch.cat([x.unsqueeze(0) for x in control_tensors])
|
||||
|
||||
self.loss_multiplier_list: List[float] = [x.loss_multiplier for x in self.file_items]
|
||||
|
||||
if any([x.clip_image_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_clip_image_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.clip_image_tensor is not None:
|
||||
base_clip_image_tensor = x.clip_image_tensor
|
||||
break
|
||||
clip_image_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_tensor is None:
|
||||
clip_image_tensors.append(torch.zeros_like(base_clip_image_tensor))
|
||||
else:
|
||||
clip_image_tensors.append(x.clip_image_tensor)
|
||||
self.clip_image_tensor = torch.cat([x.unsqueeze(0) for x in clip_image_tensors])
|
||||
|
||||
if any([x.mask_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_mask_tensor = None
|
||||
@@ -159,6 +195,23 @@ class DataLoaderBatchDTO:
|
||||
else:
|
||||
unconditional_tensor.append(x.unconditional_tensor)
|
||||
self.unconditional_tensor = torch.cat([x.unsqueeze(0) for x in unconditional_tensor])
|
||||
|
||||
if any([x.clip_image_embeds is not None for x in self.file_items]):
|
||||
self.clip_image_embeds = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_embeds is not None:
|
||||
self.clip_image_embeds.append(x.clip_image_embeds)
|
||||
else:
|
||||
raise Exception("clip_image_embeds is None for some file items")
|
||||
|
||||
if any([x.clip_image_embeds_unconditional is not None for x in self.file_items]):
|
||||
self.clip_image_embeds_unconditional = []
|
||||
for x in self.file_items:
|
||||
if x.clip_image_embeds_unconditional is not None:
|
||||
self.clip_image_embeds_unconditional.append(x.clip_image_embeds_unconditional)
|
||||
else:
|
||||
raise Exception("clip_image_embeds_unconditional is None for some file items")
|
||||
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
@@ -175,11 +228,7 @@ class DataLoaderBatchDTO:
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present
|
||||
) for x in self.file_items]
|
||||
return [x.caption for x in self.file_items]
|
||||
|
||||
def get_caption_short_list(
|
||||
self,
|
||||
@@ -187,12 +236,7 @@ class DataLoaderBatchDTO:
|
||||
to_replace_list=None,
|
||||
add_if_not_present=True
|
||||
):
|
||||
return [x.get_caption(
|
||||
trigger=trigger,
|
||||
to_replace_list=to_replace_list,
|
||||
add_if_not_present=add_if_not_present,
|
||||
short_caption=True
|
||||
) for x in self.file_items]
|
||||
return [x.caption_short for x in self.file_items]
|
||||
|
||||
def cleanup(self):
|
||||
del self.latents
|
||||
|
||||
@@ -12,6 +12,7 @@ import numpy as np
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from tqdm import tqdm
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from toolkit.basic import flush, value_map
|
||||
from toolkit.buckets import get_bucket_for_image_size, get_resolution
|
||||
@@ -27,7 +28,7 @@ from toolkit.train_tools import get_torch_dtype
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.data_loader import AiToolkitDataset
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO
|
||||
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
# def get_associated_caption_from_img_path(img_path):
|
||||
# https://demo.albumentations.ai/
|
||||
@@ -56,6 +57,30 @@ transforms_dict = {
|
||||
caption_ext_list = ['txt', 'json', 'caption']
|
||||
|
||||
|
||||
def standardize_images(images):
|
||||
"""
|
||||
Standardize the given batch of images using the specified mean and std.
|
||||
Expects values of 0 - 1
|
||||
|
||||
Args:
|
||||
images (torch.Tensor): A batch of images in the shape of (N, C, H, W),
|
||||
where N is the number of images, C is the number of channels,
|
||||
H is the height, and W is the width.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Standardized images.
|
||||
"""
|
||||
mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
std = [0.26862954, 0.26130258, 0.27577711]
|
||||
|
||||
# Define the normalization transform
|
||||
normalize = transforms.Normalize(mean=mean, std=std)
|
||||
|
||||
# Apply normalization to each image in the batch
|
||||
standardized_images = torch.stack([normalize(img) for img in images])
|
||||
|
||||
return standardized_images
|
||||
|
||||
def clean_caption(caption):
|
||||
# remove any newlines
|
||||
caption = caption.replace('\n', ', ')
|
||||
@@ -112,6 +137,13 @@ class CaptionMixin:
|
||||
prompt = self.default_prompt
|
||||
if hasattr(self, 'default_caption'):
|
||||
prompt = self.default_caption
|
||||
|
||||
# handle replacements
|
||||
replacement_list = self.dataset_config.replacements if isinstance(self.dataset_config.replacements, list) else []
|
||||
for replacement in replacement_list:
|
||||
from_string, to_string = replacement.split('|')
|
||||
prompt = prompt.replace(from_string, to_string)
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
@@ -171,7 +203,22 @@ class BucketsMixin:
|
||||
if file_item.has_point_of_interest:
|
||||
# Attempt to process the poi if we can. It wont process if the image is smaller than the resolution
|
||||
did_process_poi = file_item.setup_poi_bucket()
|
||||
if not did_process_poi:
|
||||
if self.dataset_config.square_crop:
|
||||
# we scale first so smallest size matches resolution
|
||||
scale_factor_x = resolution / width
|
||||
scale_factor_y = resolution / height
|
||||
scale_factor = max(scale_factor_x, scale_factor_y)
|
||||
file_item.scale_to_width = math.ceil(width * scale_factor)
|
||||
file_item.scale_to_height = math.ceil(height * scale_factor)
|
||||
file_item.crop_width = resolution
|
||||
file_item.crop_height = resolution
|
||||
if width > height:
|
||||
file_item.crop_x = int(file_item.scale_to_width / 2 - resolution / 2)
|
||||
file_item.crop_y = 0
|
||||
else:
|
||||
file_item.crop_x = 0
|
||||
file_item.crop_y = int(file_item.scale_to_height / 2 - resolution / 2)
|
||||
elif not did_process_poi:
|
||||
bucket_resolution = get_bucket_for_image_size(
|
||||
width, height,
|
||||
resolution=resolution,
|
||||
@@ -231,6 +278,11 @@ class CaptionProcessingDTOMixin:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.raw_caption: str = None
|
||||
self.raw_caption_short: str = None
|
||||
self.caption: str = None
|
||||
self.caption_short: str = None
|
||||
|
||||
dataset_config: DatasetConfig = kwargs.get('dataset_config', None)
|
||||
self.extra_values: List[float] = dataset_config.extra_values
|
||||
|
||||
# todo allow for loading from sd-scripts style dict
|
||||
def load_caption(self: 'FileItemDTO', caption_dict: Union[dict, None]):
|
||||
@@ -258,11 +310,15 @@ class CaptionProcessingDTOMixin:
|
||||
prompt = prompt.replace('\n', ' ')
|
||||
prompt = prompt.replace('\r', ' ')
|
||||
|
||||
prompt = json.loads(prompt)
|
||||
if 'caption' in prompt:
|
||||
prompt = prompt['caption']
|
||||
if 'caption_short' in prompt:
|
||||
short_caption = prompt['caption_short']
|
||||
prompt_json = json.loads(prompt)
|
||||
if 'caption' in prompt_json:
|
||||
prompt = prompt_json['caption']
|
||||
if 'caption_short' in prompt_json:
|
||||
short_caption = prompt_json['caption_short']
|
||||
|
||||
if 'extra_values' in prompt_json:
|
||||
self.extra_values = prompt_json['extra_values']
|
||||
|
||||
prompt = clean_caption(prompt)
|
||||
if short_caption is not None:
|
||||
short_caption = clean_caption(short_caption)
|
||||
@@ -276,6 +332,10 @@ class CaptionProcessingDTOMixin:
|
||||
self.raw_caption = prompt
|
||||
self.raw_caption_short = short_caption
|
||||
|
||||
self.caption = self.get_caption()
|
||||
if self.raw_caption_short is not None:
|
||||
self.caption_short = self.get_caption(short_caption=True)
|
||||
|
||||
def get_caption(
|
||||
self: 'FileItemDTO',
|
||||
trigger=None,
|
||||
@@ -304,27 +364,44 @@ class CaptionProcessingDTOMixin:
|
||||
# remove empty strings
|
||||
token_list = [x for x in token_list if x]
|
||||
|
||||
if self.dataset_config.shuffle_tokens:
|
||||
random.shuffle(token_list)
|
||||
|
||||
# handle token dropout
|
||||
if self.dataset_config.token_dropout_rate > 0 and not short_caption:
|
||||
new_token_list = []
|
||||
for token in token_list:
|
||||
# get a random float form 0 to 1
|
||||
rand = random.random()
|
||||
if rand > self.dataset_config.token_dropout_rate:
|
||||
# keep the token
|
||||
keep_tokens: int = self.dataset_config.keep_tokens
|
||||
for idx, token in enumerate(token_list):
|
||||
if idx < keep_tokens:
|
||||
new_token_list.append(token)
|
||||
elif self.dataset_config.token_dropout_rate >= 1.0:
|
||||
# drop the token
|
||||
pass
|
||||
else:
|
||||
# get a random float form 0 to 1
|
||||
rand = random.random()
|
||||
if rand > self.dataset_config.token_dropout_rate:
|
||||
# keep the token
|
||||
new_token_list.append(token)
|
||||
token_list = new_token_list
|
||||
|
||||
if self.dataset_config.shuffle_tokens:
|
||||
random.shuffle(token_list)
|
||||
|
||||
# join back together
|
||||
caption = ', '.join(token_list)
|
||||
caption = inject_trigger_into_prompt(caption, trigger, to_replace_list, add_if_not_present)
|
||||
# caption = inject_trigger_into_prompt(caption, trigger, to_replace_list, add_if_not_present)
|
||||
|
||||
if self.dataset_config.random_triggers and len(self.dataset_config.random_triggers) > 0:
|
||||
# add random triggers
|
||||
caption = caption + ', ' + random.choice(self.dataset_config.random_triggers)
|
||||
if self.dataset_config.random_triggers:
|
||||
num_triggers = self.dataset_config.random_triggers_max
|
||||
if num_triggers > 1:
|
||||
num_triggers = random.randint(0, num_triggers)
|
||||
|
||||
if num_triggers > 0:
|
||||
triggers = random.sample(self.dataset_config.random_triggers, num_triggers)
|
||||
caption = caption + ', ' + ', '.join(triggers)
|
||||
# add random triggers
|
||||
# for i in range(num_triggers):
|
||||
# # fastest method
|
||||
# trigger = self.dataset_config.random_triggers[int(random.random() * (len(self.dataset_config.random_triggers)))]
|
||||
# caption = caption + ', ' + trigger
|
||||
|
||||
if self.dataset_config.shuffle_tokens:
|
||||
# shuffle again
|
||||
@@ -350,6 +427,8 @@ class ImageProcessingDTOMixin:
|
||||
self.get_latent()
|
||||
if self.has_control_image:
|
||||
self.load_control_image()
|
||||
if self.has_clip_image:
|
||||
self.load_clip_image()
|
||||
if self.has_mask_image:
|
||||
self.load_mask_image()
|
||||
if self.has_unconditional:
|
||||
@@ -374,19 +453,19 @@ class ImageProcessingDTOMixin:
|
||||
w, h = img.size
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
print(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
elif h > w and self.scale_to_height < self.scale_to_width:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
print(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
if self.dataset_config.buckets:
|
||||
# scale and crop based on file item
|
||||
@@ -443,6 +522,8 @@ class ImageProcessingDTOMixin:
|
||||
if not only_load_latents:
|
||||
if self.has_control_image:
|
||||
self.load_control_image()
|
||||
if self.has_clip_image:
|
||||
self.load_clip_image()
|
||||
if self.has_mask_image:
|
||||
self.load_mask_image()
|
||||
if self.has_unconditional:
|
||||
@@ -457,9 +538,11 @@ class ControlFileItemDTOMixin:
|
||||
self.control_path: Union[str, None] = None
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
self.full_size_control_images = False
|
||||
if dataset_config.control_path is not None:
|
||||
# find the control image path
|
||||
control_path = dataset_config.control_path
|
||||
self.full_size_control_images = dataset_config.full_size_control_images
|
||||
# we are using control images
|
||||
img_path = kwargs.get('path', None)
|
||||
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
@@ -477,43 +560,254 @@ class ControlFileItemDTOMixin:
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.control_path}")
|
||||
w, h = img.size
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
elif h > w and self.scale_to_height < self.scale_to_width:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
if self.full_size_control_images:
|
||||
# we just scale them to 512x512:
|
||||
w, h = img.size
|
||||
img = img.resize((512, 512), Image.BICUBIC)
|
||||
|
||||
if self.dataset_config.buckets:
|
||||
# scale and crop based on file item
|
||||
img = img.resize((self.scale_to_width, self.scale_to_height), Image.BICUBIC)
|
||||
# img = transforms.CenterCrop((self.crop_height, self.crop_width))(img)
|
||||
# crop
|
||||
img = img.crop((
|
||||
self.crop_x,
|
||||
self.crop_y,
|
||||
self.crop_x + self.crop_width,
|
||||
self.crop_y + self.crop_height
|
||||
))
|
||||
else:
|
||||
raise Exception("Control images not supported for non-bucket datasets")
|
||||
w, h = img.size
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
elif h > w and self.scale_to_height < self.scale_to_width:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
|
||||
self.control_tensor = transforms.ToTensor()(img)
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
if self.dataset_config.buckets:
|
||||
# scale and crop based on file item
|
||||
img = img.resize((self.scale_to_width, self.scale_to_height), Image.BICUBIC)
|
||||
# img = transforms.CenterCrop((self.crop_height, self.crop_width))(img)
|
||||
# crop
|
||||
img = img.crop((
|
||||
self.crop_x,
|
||||
self.crop_y,
|
||||
self.crop_x + self.crop_width,
|
||||
self.crop_y + self.crop_height
|
||||
))
|
||||
else:
|
||||
raise Exception("Control images not supported for non-bucket datasets")
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
if self.aug_replay_spatial_transforms:
|
||||
self.control_tensor = self.augment_spatial_control(img, transform=transform)
|
||||
else:
|
||||
self.control_tensor = transform(img)
|
||||
|
||||
def cleanup_control(self: 'FileItemDTO'):
|
||||
self.control_tensor = None
|
||||
|
||||
|
||||
class ClipImageFileItemDTOMixin:
|
||||
def __init__(self: 'FileItemDTO', *args, **kwargs):
|
||||
if hasattr(super(), '__init__'):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.has_clip_image = False
|
||||
self.clip_image_path: Union[str, None] = None
|
||||
self.clip_image_tensor: Union[torch.Tensor, None] = None
|
||||
self.clip_image_embeds: Union[dict, None] = None
|
||||
self.clip_image_embeds_unconditional: Union[dict, None] = None
|
||||
self.has_clip_augmentations = False
|
||||
self.clip_image_aug_transform: Union[None, A.Compose] = None
|
||||
self.clip_image_processor: Union[None, CLIPImageProcessor] = None
|
||||
self.clip_image_encoder_path: Union[str, None] = None
|
||||
self.is_caching_clip_vision_to_disk = False
|
||||
self.is_vision_clip_cached = False
|
||||
self.clip_vision_is_quad = False
|
||||
self.clip_vision_load_device = 'cpu'
|
||||
self.clip_vision_unconditional_paths: Union[List[str], None] = None
|
||||
self._clip_vision_embeddings_path: Union[str, None] = None
|
||||
dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
if dataset_config.clip_image_path is not None:
|
||||
# copy the clip image processor so the dataloader can do it
|
||||
sd = kwargs.get('sd', None)
|
||||
if hasattr(sd.adapter, 'clip_image_processor'):
|
||||
self.clip_image_processor = sd.adapter.clip_image_processor
|
||||
# find the control image path
|
||||
clip_image_path = dataset_config.clip_image_path
|
||||
# we are using control images
|
||||
img_path = kwargs.get('path', None)
|
||||
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
file_name_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
for ext in img_ext_list:
|
||||
if os.path.exists(os.path.join(clip_image_path, file_name_no_ext + ext)):
|
||||
self.clip_image_path = os.path.join(clip_image_path, file_name_no_ext + ext)
|
||||
self.has_clip_image = True
|
||||
break
|
||||
|
||||
self.build_clip_imag_augmentation_transform()
|
||||
|
||||
def build_clip_imag_augmentation_transform(self: 'FileItemDTO'):
|
||||
if self.dataset_config.clip_image_augmentations is not None and len(self.dataset_config.clip_image_augmentations) > 0:
|
||||
self.has_clip_augmentations = True
|
||||
augmentations = [Augments(**aug) for aug in self.dataset_config.clip_image_augmentations]
|
||||
|
||||
if self.dataset_config.clip_image_shuffle_augmentations:
|
||||
random.shuffle(augmentations)
|
||||
|
||||
augmentation_list = []
|
||||
for aug in augmentations:
|
||||
# make sure method name is valid
|
||||
assert hasattr(A, aug.method_name), f"invalid augmentation method: {aug.method_name}"
|
||||
# get the method
|
||||
method = getattr(A, aug.method_name)
|
||||
# add the method to the list
|
||||
augmentation_list.append(method(**aug.params))
|
||||
|
||||
self.clip_image_aug_transform = A.Compose(augmentation_list)
|
||||
|
||||
def augment_clip_image(self: 'FileItemDTO', img: Image, transform: Union[None, transforms.Compose], ):
|
||||
if self.dataset_config.clip_image_shuffle_augmentations:
|
||||
self.build_clip_imag_augmentation_transform()
|
||||
|
||||
open_cv_image = np.array(img)
|
||||
# Convert RGB to BGR
|
||||
open_cv_image = open_cv_image[:, :, ::-1].copy()
|
||||
|
||||
if self.clip_vision_is_quad:
|
||||
# image is in a 2x2 gris. split, run augs, and recombine
|
||||
# split
|
||||
img1, img2 = np.hsplit(open_cv_image, 2)
|
||||
img1_1, img1_2 = np.vsplit(img1, 2)
|
||||
img2_1, img2_2 = np.vsplit(img2, 2)
|
||||
# apply augmentations
|
||||
img1_1 = self.clip_image_aug_transform(image=img1_1)["image"]
|
||||
img1_2 = self.clip_image_aug_transform(image=img1_2)["image"]
|
||||
img2_1 = self.clip_image_aug_transform(image=img2_1)["image"]
|
||||
img2_2 = self.clip_image_aug_transform(image=img2_2)["image"]
|
||||
# recombine
|
||||
augmented = np.vstack((np.hstack((img1_1, img1_2)), np.hstack((img2_1, img2_2))))
|
||||
|
||||
else:
|
||||
# apply augmentations
|
||||
augmented = self.clip_image_aug_transform(image=open_cv_image)["image"]
|
||||
|
||||
# convert back to RGB tensor
|
||||
augmented = cv2.cvtColor(augmented, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# convert to PIL image
|
||||
augmented = Image.fromarray(augmented)
|
||||
|
||||
augmented_tensor = transforms.ToTensor()(augmented) if transform is None else transform(augmented)
|
||||
|
||||
return augmented_tensor
|
||||
|
||||
def get_clip_vision_info_dict(self: 'FileItemDTO'):
|
||||
item = OrderedDict([
|
||||
("image_encoder_path", self.clip_image_encoder_path),
|
||||
("filename", os.path.basename(self.clip_image_path)),
|
||||
("is_quad", self.clip_vision_is_quad)
|
||||
])
|
||||
# when adding items, do it after so we dont change old latents
|
||||
if self.flip_x:
|
||||
item["flip_x"] = True
|
||||
if self.flip_y:
|
||||
item["flip_y"] = True
|
||||
return item
|
||||
def get_clip_vision_embeddings_path(self: 'FileItemDTO', recalculate=False):
|
||||
if self._clip_vision_embeddings_path is not None and not recalculate:
|
||||
return self._clip_vision_embeddings_path
|
||||
else:
|
||||
# we store latents in a folder in same path as image called _latent_cache
|
||||
img_dir = os.path.dirname(self.clip_image_path)
|
||||
latent_dir = os.path.join(img_dir, '_clip_vision_cache')
|
||||
hash_dict = self.get_clip_vision_info_dict()
|
||||
filename_no_ext = os.path.splitext(os.path.basename(self.clip_image_path))[0]
|
||||
# get base64 hash of md5 checksum of hash_dict
|
||||
hash_input = json.dumps(hash_dict, sort_keys=True).encode('utf-8')
|
||||
hash_str = base64.urlsafe_b64encode(hashlib.md5(hash_input).digest()).decode('ascii')
|
||||
hash_str = hash_str.replace('=', '')
|
||||
self._clip_vision_embeddings_path = os.path.join(latent_dir, f'{filename_no_ext}_{hash_str}.safetensors')
|
||||
|
||||
return self._clip_vision_embeddings_path
|
||||
|
||||
def load_clip_image(self: 'FileItemDTO'):
|
||||
if self.is_vision_clip_cached:
|
||||
self.clip_image_embeds = load_file(self.get_clip_vision_embeddings_path())
|
||||
|
||||
# get a random unconditional image
|
||||
if self.clip_vision_unconditional_paths is not None:
|
||||
unconditional_path = random.choice(self.clip_vision_unconditional_paths)
|
||||
self.clip_image_embeds_unconditional = load_file(unconditional_path)
|
||||
|
||||
return
|
||||
try:
|
||||
img = Image.open(self.clip_image_path).convert('RGB')
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
# make a random noise image
|
||||
img = Image.new('RGB', (self.dataset_config.resolution, self.dataset_config.resolution))
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.clip_image_path}")
|
||||
|
||||
img = img.convert('RGB')
|
||||
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
if img.width != img.height:
|
||||
min_size = min(img.width, img.height)
|
||||
if self.dataset_config.square_crop:
|
||||
# center crop to a square
|
||||
img = transforms.CenterCrop(min_size)(img)
|
||||
else:
|
||||
# image must be square. If it is not, we will resize/squish it so it is, that way we don't crop out data
|
||||
# resize to the smallest dimension
|
||||
img = img.resize((min_size, min_size), Image.BICUBIC)
|
||||
|
||||
if self.has_clip_augmentations:
|
||||
self.clip_image_tensor = self.augment_clip_image(img, transform=None)
|
||||
else:
|
||||
self.clip_image_tensor = transforms.ToTensor()(img)
|
||||
|
||||
# random crop
|
||||
# if self.dataset_config.clip_image_random_crop:
|
||||
# # crop up to 20% on all sides. Keep is square
|
||||
# crop_percent = random.randint(0, 20) / 100
|
||||
# crop_width = int(self.clip_image_tensor.shape[2] * crop_percent)
|
||||
# crop_height = int(self.clip_image_tensor.shape[1] * crop_percent)
|
||||
# crop_left = random.randint(0, crop_width)
|
||||
# crop_top = random.randint(0, crop_height)
|
||||
# crop_right = self.clip_image_tensor.shape[2] - crop_width - crop_left
|
||||
# crop_bottom = self.clip_image_tensor.shape[1] - crop_height - crop_top
|
||||
# if len(self.clip_image_tensor.shape) == 3:
|
||||
# self.clip_image_tensor = self.clip_image_tensor[:, crop_top:-crop_bottom, crop_left:-crop_right]
|
||||
# elif len(self.clip_image_tensor.shape) == 4:
|
||||
# self.clip_image_tensor = self.clip_image_tensor[:, :, crop_top:-crop_bottom, crop_left:-crop_right]
|
||||
|
||||
if self.clip_image_processor is not None:
|
||||
# run it
|
||||
tensors_0_1 = self.clip_image_tensor.to(dtype=torch.float16)
|
||||
clip_out = self.clip_image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
self.clip_image_tensor = clip_out.squeeze(0).clone().detach()
|
||||
|
||||
def cleanup_clip_image(self: 'FileItemDTO'):
|
||||
self.clip_image_tensor = None
|
||||
self.clip_image_embeds = None
|
||||
|
||||
|
||||
|
||||
|
||||
class AugmentationFileItemDTOMixin:
|
||||
def __init__(self: 'FileItemDTO', *args, **kwargs):
|
||||
if hasattr(super(), '__init__'):
|
||||
@@ -522,6 +816,8 @@ class AugmentationFileItemDTOMixin:
|
||||
self.unaugmented_tensor: Union[torch.Tensor, None] = None
|
||||
# self.augmentations: Union[None, List[Augments]] = None
|
||||
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
self.aug_transform: Union[None, A.Compose] = None
|
||||
self.aug_replay_spatial_transforms = None
|
||||
self.build_augmentation_transform()
|
||||
|
||||
def build_augmentation_transform(self: 'FileItemDTO'):
|
||||
@@ -541,7 +837,8 @@ class AugmentationFileItemDTOMixin:
|
||||
# add the method to the list
|
||||
augmentation_list.append(method(**aug.params))
|
||||
|
||||
self.aug_transform = A.Compose(augmentation_list)
|
||||
# add additional targets so we can augment the control image
|
||||
self.aug_transform = A.ReplayCompose(augmentation_list, additional_targets={'image2': 'image'})
|
||||
|
||||
def augment_image(self: 'FileItemDTO', img: Image, transform: Union[None, transforms.Compose], ):
|
||||
|
||||
@@ -557,7 +854,18 @@ class AugmentationFileItemDTOMixin:
|
||||
open_cv_image = open_cv_image[:, :, ::-1].copy()
|
||||
|
||||
# apply augmentations
|
||||
augmented = self.aug_transform(image=open_cv_image)["image"]
|
||||
transformed = self.aug_transform(image=open_cv_image)
|
||||
augmented = transformed["image"]
|
||||
|
||||
# save just the spatial transforms for controls and masks
|
||||
augmented_params = transformed["replay"]
|
||||
spatial_transforms = ['Rotate', 'Flip', 'HorizontalFlip', 'VerticalFlip', 'Resize', 'Crop', 'RandomCrop',
|
||||
'ElasticTransform', 'GridDistortion', 'OpticalDistortion']
|
||||
# only store the spatial transforms
|
||||
augmented_params['transforms'] = [t for t in augmented_params['transforms'] if t['__class_fullname__'].split('.')[-1] in spatial_transforms]
|
||||
|
||||
if self.dataset_config.replay_transforms:
|
||||
self.aug_replay_spatial_transforms = augmented_params
|
||||
|
||||
# convert back to RGB tensor
|
||||
augmented = cv2.cvtColor(augmented, cv2.COLOR_BGR2RGB)
|
||||
@@ -569,6 +877,38 @@ class AugmentationFileItemDTOMixin:
|
||||
|
||||
return augmented_tensor
|
||||
|
||||
# augment control images spatially consistent with transforms done to the main image
|
||||
def augment_spatial_control(self: 'FileItemDTO', img: Image, transform: Union[None, transforms.Compose] ):
|
||||
if self.aug_replay_spatial_transforms is None:
|
||||
# no transforms
|
||||
return transform(img)
|
||||
|
||||
# save colorspace to convert back to
|
||||
colorspace = img.mode
|
||||
|
||||
# convert to rgb
|
||||
img = img.convert('RGB')
|
||||
|
||||
open_cv_image = np.array(img)
|
||||
# Convert RGB to BGR
|
||||
open_cv_image = open_cv_image[:, :, ::-1].copy()
|
||||
|
||||
# Replay transforms
|
||||
transformed = A.ReplayCompose.replay(self.aug_replay_spatial_transforms, image=open_cv_image)
|
||||
augmented = transformed["image"]
|
||||
|
||||
# convert back to RGB tensor
|
||||
augmented = cv2.cvtColor(augmented, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# convert to PIL image
|
||||
augmented = Image.fromarray(augmented)
|
||||
|
||||
# convert back to original colorspace
|
||||
augmented = augmented.convert(colorspace)
|
||||
|
||||
augmented_tensor = transforms.ToTensor()(augmented) if transform is None else transform(augmented)
|
||||
return augmented_tensor
|
||||
|
||||
def cleanup_control(self: 'FileItemDTO'):
|
||||
self.unaugmented_tensor = None
|
||||
|
||||
@@ -620,21 +960,31 @@ class MaskFileItemDTOMixin:
|
||||
if self.dataset_config.invert_mask:
|
||||
img = ImageOps.invert(img)
|
||||
w, h = img.size
|
||||
fix_size = False
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
print(f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
fix_size = True
|
||||
elif h > w and self.scale_to_height < self.scale_to_width:
|
||||
# throw error, they should match
|
||||
raise ValueError(
|
||||
f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
print(f"unexpected values: w={w}, h={h}, file_item.scale_to_width={self.scale_to_width}, file_item.scale_to_height={self.scale_to_height}, file_item.path={self.path}")
|
||||
fix_size = True
|
||||
|
||||
if fix_size:
|
||||
# swap all the sizes
|
||||
self.scale_to_width, self.scale_to_height = self.scale_to_height, self.scale_to_width
|
||||
self.crop_width, self.crop_height = self.crop_height, self.crop_width
|
||||
self.crop_x, self.crop_y = self.crop_y, self.crop_x
|
||||
|
||||
|
||||
|
||||
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
# randomly apply a blur up to 0.5% of the size of the min (width, height)
|
||||
min_size = min(img.width, img.height)
|
||||
@@ -658,7 +1008,13 @@ class MaskFileItemDTOMixin:
|
||||
else:
|
||||
raise Exception("Mask images not supported for non-bucket datasets")
|
||||
|
||||
self.mask_tensor = transforms.ToTensor()(img)
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
if self.aug_replay_spatial_transforms:
|
||||
self.mask_tensor = self.augment_spatial_control(img, transform=transform)
|
||||
else:
|
||||
self.mask_tensor = transform(img)
|
||||
self.mask_tensor = value_map(self.mask_tensor, 0, 1.0, self.mask_min_value, 1.0)
|
||||
# convert to grayscale
|
||||
|
||||
@@ -674,12 +1030,7 @@ class UnconditionalFileItemDTOMixin:
|
||||
self.unconditional_path: Union[str, None] = None
|
||||
self.unconditional_tensor: Union[torch.Tensor, None] = None
|
||||
self.unconditional_latent: Union[torch.Tensor, None] = None
|
||||
self.unconditional_transforms = transforms.Compose(
|
||||
[
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5], [0.5]),
|
||||
]
|
||||
)
|
||||
self.unconditional_transforms = self.dataloader_transforms
|
||||
dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
|
||||
if dataset_config.unconditional_path is not None:
|
||||
@@ -714,10 +1065,10 @@ class UnconditionalFileItemDTOMixin:
|
||||
|
||||
if self.flip_x:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
if self.dataset_config.buckets:
|
||||
# scale and crop based on file item
|
||||
@@ -733,7 +1084,10 @@ class UnconditionalFileItemDTOMixin:
|
||||
else:
|
||||
raise Exception("Unconditional images are not supported for non-bucket datasets")
|
||||
|
||||
self.unconditional_tensor = self.unconditional_transforms(img)
|
||||
if self.aug_replay_spatial_transforms:
|
||||
self.unconditional_tensor = self.augment_spatial_control(img, transform=self.unconditional_transforms)
|
||||
else:
|
||||
self.unconditional_tensor = self.unconditional_transforms(img)
|
||||
|
||||
def cleanup_unconditional(self: 'FileItemDTO'):
|
||||
self.unconditional_tensor = None
|
||||
@@ -848,8 +1202,14 @@ class PoiFileItemDTOMixin:
|
||||
crop_bottom = initial_height
|
||||
|
||||
poi_height = crop_bottom - poi_y
|
||||
# now we have our random crop, but it may be smaller than resolution. Check and expand if needed
|
||||
current_resolution = get_resolution(poi_width, poi_height)
|
||||
try:
|
||||
# now we have our random crop, but it may be smaller than resolution. Check and expand if needed
|
||||
current_resolution = get_resolution(poi_width, poi_height)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error getting resolution: {self.path}")
|
||||
raise e
|
||||
return False
|
||||
if current_resolution >= self.dataset_config.resolution:
|
||||
# We can break now
|
||||
break
|
||||
@@ -874,13 +1234,17 @@ class PoiFileItemDTOMixin:
|
||||
# Use the maximum of the scale factors to ensure both dimensions are scaled above the bucket resolution
|
||||
max_scale_factor = max(width_scale_factor, height_scale_factor)
|
||||
|
||||
self.scale_to_width = int(initial_width * max_scale_factor)
|
||||
self.scale_to_height = int(initial_height * max_scale_factor)
|
||||
self.scale_to_width = math.ceil(initial_width * max_scale_factor)
|
||||
self.scale_to_height = math.ceil(initial_height * max_scale_factor)
|
||||
self.crop_width = bucket_resolution['width']
|
||||
self.crop_height = bucket_resolution['height']
|
||||
self.crop_x = int(poi_x * max_scale_factor)
|
||||
self.crop_y = int(poi_y * max_scale_factor)
|
||||
|
||||
if self.scale_to_width < self.crop_x + self.crop_width or self.scale_to_height < self.crop_y + self.crop_height:
|
||||
# todo look into this. This still happens sometimes
|
||||
print('size mismatch')
|
||||
|
||||
return True
|
||||
|
||||
|
||||
@@ -989,7 +1353,17 @@ class LatentCachingMixin:
|
||||
i = 0
|
||||
for file_item in tqdm(self.file_list, desc=f'Caching latents{" to disk" if to_disk else ""}'):
|
||||
# set latent space version
|
||||
if self.sd.is_xl:
|
||||
if self.sd.model_config.latent_space_version is not None:
|
||||
file_item.latent_space_version = self.sd.model_config.latent_space_version
|
||||
elif self.sd.is_xl:
|
||||
file_item.latent_space_version = 'sdxl'
|
||||
elif self.sd.is_v3:
|
||||
file_item.latent_space_version = 'sd3'
|
||||
elif self.sd.is_auraflow:
|
||||
file_item.latent_space_version = 'sdxl'
|
||||
elif self.sd.is_flux:
|
||||
file_item.latent_space_version = 'flux1'
|
||||
elif self.sd.model_config.is_pixart_sigma:
|
||||
file_item.latent_space_version = 'sdxl'
|
||||
else:
|
||||
file_item.latent_space_version = 'sd1'
|
||||
@@ -1011,8 +1385,13 @@ class LatentCachingMixin:
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
# add batch dimension
|
||||
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
latent = self.sd.encode_images(imgs).squeeze(0)
|
||||
try:
|
||||
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
latent = self.sd.encode_images(imgs).squeeze(0)
|
||||
except Exception as e:
|
||||
print(f"Error processing image: {file_item.path}")
|
||||
print(f"Error: {str(e)}")
|
||||
raise e
|
||||
# save_latent
|
||||
if to_disk:
|
||||
state_dict = OrderedDict([
|
||||
@@ -1031,7 +1410,7 @@ class LatentCachingMixin:
|
||||
del latent
|
||||
del file_item.tensor
|
||||
|
||||
flush(garbage_collect=False)
|
||||
# flush(garbage_collect=False)
|
||||
file_item.is_latent_cached = True
|
||||
i += 1
|
||||
# flush every 100
|
||||
@@ -1040,3 +1419,176 @@ class LatentCachingMixin:
|
||||
|
||||
# restore device state
|
||||
self.sd.restore_device_state()
|
||||
|
||||
|
||||
class CLIPCachingMixin:
|
||||
def __init__(self: 'AiToolkitDataset', **kwargs):
|
||||
# if we have super, call it
|
||||
if hasattr(super(), '__init__'):
|
||||
super().__init__(**kwargs)
|
||||
self.clip_vision_num_unconditional_cache = 20
|
||||
self.clip_vision_unconditional_cache = []
|
||||
|
||||
def cache_clip_vision_to_disk(self: 'AiToolkitDataset'):
|
||||
if not self.is_caching_clip_vision_to_disk:
|
||||
return
|
||||
with torch.no_grad():
|
||||
print(f"Caching clip vision for {self.dataset_path}")
|
||||
|
||||
print(" - Saving clip to disk")
|
||||
# move sd items to cpu except for vae
|
||||
self.sd.set_device_state_preset('cache_clip')
|
||||
|
||||
# make sure the adapter has attributes
|
||||
if self.sd.adapter is None:
|
||||
raise Exception("Error: must have an adapter to cache clip vision to disk")
|
||||
|
||||
clip_image_processor: CLIPImageProcessor = None
|
||||
if hasattr(self.sd.adapter, 'clip_image_processor'):
|
||||
clip_image_processor = self.sd.adapter.clip_image_processor
|
||||
|
||||
if clip_image_processor is None:
|
||||
raise Exception("Error: must have a clip image processor to cache clip vision to disk")
|
||||
|
||||
vision_encoder: CLIPVisionModelWithProjection = None
|
||||
if hasattr(self.sd.adapter, 'image_encoder'):
|
||||
vision_encoder = self.sd.adapter.image_encoder
|
||||
if hasattr(self.sd.adapter, 'vision_encoder'):
|
||||
vision_encoder = self.sd.adapter.vision_encoder
|
||||
|
||||
if vision_encoder is None:
|
||||
raise Exception("Error: must have a vision encoder to cache clip vision to disk")
|
||||
|
||||
# move vision encoder to device
|
||||
vision_encoder.to(self.sd.device)
|
||||
|
||||
is_quad = self.sd.adapter.config.quad_image
|
||||
image_encoder_path = self.sd.adapter.config.image_encoder_path
|
||||
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
if hasattr(self.sd.adapter, 'clip_noise_zero') and self.sd.adapter.clip_noise_zero:
|
||||
# just to do this, we did :)
|
||||
# need more samples as it is random noise
|
||||
self.clip_vision_num_unconditional_cache = self.clip_vision_num_unconditional_cache
|
||||
else:
|
||||
# only need one since it doesnt change
|
||||
self.clip_vision_num_unconditional_cache = 1
|
||||
|
||||
# cache unconditionals
|
||||
print(f" - Caching {self.clip_vision_num_unconditional_cache} unconditional clip vision to disk")
|
||||
clip_vision_cache_path = os.path.join(self.dataset_config.clip_image_path, '_clip_vision_cache')
|
||||
|
||||
unconditional_paths = []
|
||||
|
||||
is_noise_zero = hasattr(self.sd.adapter, 'clip_noise_zero') and self.sd.adapter.clip_noise_zero
|
||||
|
||||
for i in range(self.clip_vision_num_unconditional_cache):
|
||||
hash_dict = OrderedDict([
|
||||
("image_encoder_path", image_encoder_path),
|
||||
("is_quad", is_quad),
|
||||
("is_noise_zero", is_noise_zero),
|
||||
])
|
||||
# get base64 hash of md5 checksum of hash_dict
|
||||
hash_input = json.dumps(hash_dict, sort_keys=True).encode('utf-8')
|
||||
hash_str = base64.urlsafe_b64encode(hashlib.md5(hash_input).digest()).decode('ascii')
|
||||
hash_str = hash_str.replace('=', '')
|
||||
|
||||
uncond_path = os.path.join(clip_vision_cache_path, f'uncond_{hash_str}_{i}.safetensors')
|
||||
if os.path.exists(uncond_path):
|
||||
# skip it
|
||||
unconditional_paths.append(uncond_path)
|
||||
continue
|
||||
|
||||
# generate a random image
|
||||
img_shape = (1, 3, self.sd.adapter.input_size, self.sd.adapter.input_size)
|
||||
if is_noise_zero:
|
||||
tensors_0_1 = torch.rand(img_shape).to(device, dtype=torch.float32)
|
||||
else:
|
||||
tensors_0_1 = torch.zeros(img_shape).to(device, dtype=torch.float32)
|
||||
clip_image = clip_image_processor(
|
||||
images=tensors_0_1,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
).pixel_values
|
||||
|
||||
if is_quad:
|
||||
# split the 4x4 grid and stack on batch
|
||||
ci1, ci2 = clip_image.chunk(2, dim=2)
|
||||
ci1, ci3 = ci1.chunk(2, dim=3)
|
||||
ci2, ci4 = ci2.chunk(2, dim=3)
|
||||
clip_image = torch.cat([ci1, ci2, ci3, ci4], dim=0).detach()
|
||||
|
||||
clip_output = vision_encoder(
|
||||
clip_image.to(device, dtype=dtype),
|
||||
output_hidden_states=True
|
||||
)
|
||||
# make state_dict ['last_hidden_state', 'image_embeds', 'penultimate_hidden_states']
|
||||
state_dict = OrderedDict([
|
||||
('image_embeds', clip_output.image_embeds.clone().detach().cpu()),
|
||||
('last_hidden_state', clip_output.hidden_states[-1].clone().detach().cpu()),
|
||||
('penultimate_hidden_states', clip_output.hidden_states[-2].clone().detach().cpu()),
|
||||
])
|
||||
|
||||
os.makedirs(os.path.dirname(uncond_path), exist_ok=True)
|
||||
save_file(state_dict, uncond_path)
|
||||
unconditional_paths.append(uncond_path)
|
||||
|
||||
self.clip_vision_unconditional_cache = unconditional_paths
|
||||
|
||||
# use tqdm to show progress
|
||||
i = 0
|
||||
for file_item in tqdm(self.file_list, desc=f'Caching clip vision to disk'):
|
||||
file_item.is_caching_clip_vision_to_disk = True
|
||||
file_item.clip_vision_load_device = self.sd.device
|
||||
file_item.clip_vision_is_quad = is_quad
|
||||
file_item.clip_image_encoder_path = image_encoder_path
|
||||
file_item.clip_vision_unconditional_paths = unconditional_paths
|
||||
if file_item.has_clip_augmentations:
|
||||
raise Exception("Error: clip vision caching is not supported with clip augmentations")
|
||||
|
||||
embedding_path = file_item.get_clip_vision_embeddings_path(recalculate=True)
|
||||
# check if it is saved to disk already
|
||||
if not os.path.exists(embedding_path):
|
||||
# load the image first
|
||||
file_item.load_clip_image()
|
||||
# add batch dimension
|
||||
clip_image = file_item.clip_image_tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
|
||||
if is_quad:
|
||||
# split the 4x4 grid and stack on batch
|
||||
ci1, ci2 = clip_image.chunk(2, dim=2)
|
||||
ci1, ci3 = ci1.chunk(2, dim=3)
|
||||
ci2, ci4 = ci2.chunk(2, dim=3)
|
||||
clip_image = torch.cat([ci1, ci2, ci3, ci4], dim=0).detach()
|
||||
|
||||
clip_output = vision_encoder(
|
||||
clip_image.to(device, dtype=dtype),
|
||||
output_hidden_states=True
|
||||
)
|
||||
|
||||
# make state_dict ['last_hidden_state', 'image_embeds', 'penultimate_hidden_states']
|
||||
state_dict = OrderedDict([
|
||||
('image_embeds', clip_output.image_embeds.clone().detach().cpu()),
|
||||
('last_hidden_state', clip_output.hidden_states[-1].clone().detach().cpu()),
|
||||
('penultimate_hidden_states', clip_output.hidden_states[-2].clone().detach().cpu()),
|
||||
])
|
||||
# metadata
|
||||
meta = get_meta_for_safetensors(file_item.get_clip_vision_info_dict())
|
||||
os.makedirs(os.path.dirname(embedding_path), exist_ok=True)
|
||||
save_file(state_dict, embedding_path, metadata=meta)
|
||||
|
||||
del clip_image
|
||||
del clip_output
|
||||
del file_item.clip_image_tensor
|
||||
|
||||
# flush(garbage_collect=False)
|
||||
file_item.is_vision_clip_cached = True
|
||||
i += 1
|
||||
# flush every 100
|
||||
# if i % 100 == 0:
|
||||
# flush()
|
||||
|
||||
# restore device state
|
||||
self.sd.restore_device_state()
|
||||
|
||||
324
toolkit/ema.py
Normal file
324
toolkit/ema.py
Normal file
@@ -0,0 +1,324 @@
|
||||
from __future__ import division
|
||||
from __future__ import unicode_literals
|
||||
|
||||
from typing import Iterable, Optional
|
||||
import weakref
|
||||
import copy
|
||||
import contextlib
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
# Partially based on:
|
||||
# https://github.com/tensorflow/tensorflow/blob/r1.13/tensorflow/python/training/moving_averages.py
|
||||
class ExponentialMovingAverage:
|
||||
"""
|
||||
Maintains (exponential) moving average of a set of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter` (typically from
|
||||
`model.parameters()`).
|
||||
Note that EMA is computed on *all* provided parameters,
|
||||
regardless of whether or not they have `requires_grad = True`;
|
||||
this allows a single EMA object to be consistantly used even
|
||||
if which parameters are trainable changes step to step.
|
||||
|
||||
If you want to some parameters in the EMA, do not pass them
|
||||
to the object in the first place. For example:
|
||||
|
||||
ExponentialMovingAverage(
|
||||
parameters=[p for p in model.parameters() if p.requires_grad],
|
||||
decay=0.9
|
||||
)
|
||||
|
||||
will ignore parameters that do not require grad.
|
||||
|
||||
decay: The exponential decay.
|
||||
|
||||
use_num_updates: Whether to use number of updates when computing
|
||||
averages.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter] = None,
|
||||
decay: float = 0.995,
|
||||
use_num_updates: bool = True,
|
||||
# feeds back the decat to the parameter
|
||||
use_feedback: bool = False
|
||||
):
|
||||
if parameters is None:
|
||||
raise ValueError("parameters must be provided")
|
||||
if decay < 0.0 or decay > 1.0:
|
||||
raise ValueError('Decay must be between 0 and 1')
|
||||
self.decay = decay
|
||||
self.num_updates = 0 if use_num_updates else None
|
||||
self.use_feedback = use_feedback
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [
|
||||
p.clone().detach()
|
||||
for p in parameters
|
||||
]
|
||||
self.collected_params = None
|
||||
self._is_train_mode = True
|
||||
# By maintaining only a weakref to each parameter,
|
||||
# we maintain the old GC behaviour of ExponentialMovingAverage:
|
||||
# if the model goes out of scope but the ExponentialMovingAverage
|
||||
# is kept, no references to the model or its parameters will be
|
||||
# maintained, and the model will be cleaned up.
|
||||
self._params_refs = [weakref.ref(p) for p in parameters]
|
||||
|
||||
def _get_parameters(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]]
|
||||
) -> Iterable[torch.nn.Parameter]:
|
||||
if parameters is None:
|
||||
parameters = [p() for p in self._params_refs]
|
||||
if any(p is None for p in parameters):
|
||||
raise ValueError(
|
||||
"(One of) the parameters with which this "
|
||||
"ExponentialMovingAverage "
|
||||
"was initialized no longer exists (was garbage collected);"
|
||||
" please either provide `parameters` explicitly or keep "
|
||||
"the model to which they belong from being garbage "
|
||||
"collected."
|
||||
)
|
||||
return parameters
|
||||
else:
|
||||
parameters = list(parameters)
|
||||
if len(parameters) != len(self.shadow_params):
|
||||
raise ValueError(
|
||||
"Number of parameters passed as argument is different "
|
||||
"from number of shadow parameters maintained by this "
|
||||
"ExponentialMovingAverage"
|
||||
)
|
||||
return parameters
|
||||
|
||||
def update(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Update currently maintained parameters.
|
||||
|
||||
Call this every time the parameters are updated, such as the result of
|
||||
the `optimizer.step()` call.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; usually the same set of
|
||||
parameters used to initialize this object. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
decay = self.decay
|
||||
if self.num_updates is not None:
|
||||
self.num_updates += 1
|
||||
decay = min(
|
||||
decay,
|
||||
(1 + self.num_updates) / (10 + self.num_updates)
|
||||
)
|
||||
one_minus_decay = 1.0 - decay
|
||||
with torch.no_grad():
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
tmp = (s_param - param)
|
||||
# tmp will be a new tensor so we can do in-place
|
||||
tmp.mul_(one_minus_decay)
|
||||
s_param.sub_(tmp)
|
||||
|
||||
if self.use_feedback:
|
||||
param.add_(tmp)
|
||||
|
||||
def copy_to(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Copy current averaged parameters into given collection of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored moving averages. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
param.data.copy_(s_param.data)
|
||||
|
||||
def store(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Save the current parameters for restoring later.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored. If `None`, the parameters of with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
self.collected_params = [
|
||||
param.clone()
|
||||
for param in parameters
|
||||
]
|
||||
|
||||
def restore(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
) -> None:
|
||||
"""
|
||||
Restore the parameters stored with the `store` method.
|
||||
Useful to validate the model with EMA parameters without affecting the
|
||||
original optimization process. Store the parameters before the
|
||||
`copy_to` method. After validation (or model saving), use this to
|
||||
restore the former parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
if self.collected_params is None:
|
||||
raise RuntimeError(
|
||||
"This ExponentialMovingAverage has no `store()`ed weights "
|
||||
"to `restore()`"
|
||||
)
|
||||
parameters = self._get_parameters(parameters)
|
||||
for c_param, param in zip(self.collected_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
|
||||
@contextlib.contextmanager
|
||||
def average_parameters(
|
||||
self,
|
||||
parameters: Optional[Iterable[torch.nn.Parameter]] = None
|
||||
):
|
||||
r"""
|
||||
Context manager for validation/inference with averaged parameters.
|
||||
|
||||
Equivalent to:
|
||||
|
||||
ema.store()
|
||||
ema.copy_to()
|
||||
try:
|
||||
...
|
||||
finally:
|
||||
ema.restore()
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters. If `None`, the
|
||||
parameters with which this `ExponentialMovingAverage` was
|
||||
initialized will be used.
|
||||
"""
|
||||
parameters = self._get_parameters(parameters)
|
||||
self.store(parameters)
|
||||
self.copy_to(parameters)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self.restore(parameters)
|
||||
|
||||
def to(self, device=None, dtype=None) -> None:
|
||||
r"""Move internal buffers of the ExponentialMovingAverage to `device`.
|
||||
|
||||
Args:
|
||||
device: like `device` argument to `torch.Tensor.to`
|
||||
"""
|
||||
# .to() on the tensors handles None correctly
|
||||
self.shadow_params = [
|
||||
p.to(device=device, dtype=dtype)
|
||||
if p.is_floating_point()
|
||||
else p.to(device=device)
|
||||
for p in self.shadow_params
|
||||
]
|
||||
if self.collected_params is not None:
|
||||
self.collected_params = [
|
||||
p.to(device=device, dtype=dtype)
|
||||
if p.is_floating_point()
|
||||
else p.to(device=device)
|
||||
for p in self.collected_params
|
||||
]
|
||||
return
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
r"""Returns the state of the ExponentialMovingAverage as a dict."""
|
||||
# Following PyTorch conventions, references to tensors are returned:
|
||||
# "returns a reference to the state and not its copy!" -
|
||||
# https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict
|
||||
return {
|
||||
"decay": self.decay,
|
||||
"num_updates": self.num_updates,
|
||||
"shadow_params": self.shadow_params,
|
||||
"collected_params": self.collected_params
|
||||
}
|
||||
|
||||
def load_state_dict(self, state_dict: dict) -> None:
|
||||
r"""Loads the ExponentialMovingAverage state.
|
||||
|
||||
Args:
|
||||
state_dict (dict): EMA state. Should be an object returned
|
||||
from a call to :meth:`state_dict`.
|
||||
"""
|
||||
# deepcopy, to be consistent with module API
|
||||
state_dict = copy.deepcopy(state_dict)
|
||||
self.decay = state_dict["decay"]
|
||||
if self.decay < 0.0 or self.decay > 1.0:
|
||||
raise ValueError('Decay must be between 0 and 1')
|
||||
self.num_updates = state_dict["num_updates"]
|
||||
assert self.num_updates is None or isinstance(self.num_updates, int), \
|
||||
"Invalid num_updates"
|
||||
|
||||
self.shadow_params = state_dict["shadow_params"]
|
||||
assert isinstance(self.shadow_params, list), \
|
||||
"shadow_params must be a list"
|
||||
assert all(
|
||||
isinstance(p, torch.Tensor) for p in self.shadow_params
|
||||
), "shadow_params must all be Tensors"
|
||||
|
||||
self.collected_params = state_dict["collected_params"]
|
||||
if self.collected_params is not None:
|
||||
assert isinstance(self.collected_params, list), \
|
||||
"collected_params must be a list"
|
||||
assert all(
|
||||
isinstance(p, torch.Tensor) for p in self.collected_params
|
||||
), "collected_params must all be Tensors"
|
||||
assert len(self.collected_params) == len(self.shadow_params), \
|
||||
"collected_params and shadow_params had different lengths"
|
||||
|
||||
if len(self.shadow_params) == len(self._params_refs):
|
||||
# Consistant with torch.optim.Optimizer, cast things to consistant
|
||||
# device and dtype with the parameters
|
||||
params = [p() for p in self._params_refs]
|
||||
# If parameters have been garbage collected, just load the state
|
||||
# we were given without change.
|
||||
if not any(p is None for p in params):
|
||||
# ^ parameter references are still good
|
||||
for i, p in enumerate(params):
|
||||
self.shadow_params[i] = self.shadow_params[i].to(
|
||||
device=p.device, dtype=p.dtype
|
||||
)
|
||||
if self.collected_params is not None:
|
||||
self.collected_params[i] = self.collected_params[i].to(
|
||||
device=p.device, dtype=p.dtype
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Tried to `load_state_dict()` with the wrong number of "
|
||||
"parameters in the saved state."
|
||||
)
|
||||
|
||||
def eval(self):
|
||||
if self._is_train_mode:
|
||||
with torch.no_grad():
|
||||
self.store()
|
||||
self.copy_to()
|
||||
self._is_train_mode = False
|
||||
|
||||
def train(self):
|
||||
if not self._is_train_mode:
|
||||
with torch.no_grad():
|
||||
self.restore()
|
||||
self._is_train_mode = True
|
||||
@@ -86,18 +86,19 @@ class Embedding:
|
||||
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
|
||||
|
||||
def restore_embeddings(self):
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[
|
||||
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
with torch.no_grad():
|
||||
with torch.no_grad():
|
||||
# Let's make sure we don't update any embedding weights besides the newly added token
|
||||
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
|
||||
self.tokenizer_list,
|
||||
self.orig_embeds_params,
|
||||
self.placeholder_token_ids):
|
||||
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||
index_no_updates[ min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
|
||||
text_encoder.get_input_embeddings().weight[
|
||||
index_no_updates
|
||||
] = orig_embeds[index_no_updates]
|
||||
weight = text_encoder.get_input_embeddings().weight
|
||||
pass
|
||||
|
||||
def get_trainable_params(self):
|
||||
params = []
|
||||
|
||||
693
toolkit/guidance.py
Normal file
693
toolkit/guidance.py
Normal file
@@ -0,0 +1,693 @@
|
||||
import torch
|
||||
from typing import Literal, Optional
|
||||
|
||||
from toolkit.basic import value_map
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
GuidanceType = Literal["targeted", "polarity", "targeted_polarity", "direct"]
|
||||
|
||||
DIFFERENTIAL_SCALER = 0.2
|
||||
|
||||
|
||||
# DIFFERENTIAL_SCALER = 0.25
|
||||
|
||||
|
||||
def get_differential_mask(
|
||||
conditional_latents: torch.Tensor,
|
||||
unconditional_latents: torch.Tensor,
|
||||
threshold: float = 0.2,
|
||||
gradient: bool = False,
|
||||
):
|
||||
# make a differential mask
|
||||
differential_mask = torch.abs(conditional_latents - unconditional_latents)
|
||||
max_differential = \
|
||||
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
differential_scaler = 1.0 / max_differential
|
||||
differential_mask = differential_mask * differential_scaler
|
||||
|
||||
if gradient:
|
||||
# wew need to scale it to 0-1
|
||||
# differential_mask = differential_mask - differential_mask.min()
|
||||
# differential_mask = differential_mask / differential_mask.max()
|
||||
# add 0.2 threshold to both sides and clip
|
||||
differential_mask = value_map(
|
||||
differential_mask,
|
||||
differential_mask.min(),
|
||||
differential_mask.max(),
|
||||
0 - threshold,
|
||||
1 + threshold
|
||||
)
|
||||
differential_mask = torch.clamp(differential_mask, 0.0, 1.0)
|
||||
else:
|
||||
|
||||
# make everything less than 0.2 be 0.0 and everything else be 1.0
|
||||
differential_mask = torch.where(
|
||||
differential_mask < threshold,
|
||||
torch.zeros_like(differential_mask),
|
||||
torch.ones_like(differential_mask)
|
||||
)
|
||||
return differential_mask
|
||||
|
||||
|
||||
def get_targeted_polarity_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
# inputs_abs_mean = torch.abs(conditional_latents).mean(dim=[1, 2, 3], keepdim=True)
|
||||
# noise_abs_mean = torch.abs(noise).mean(dim=[1, 2, 3], keepdim=True)
|
||||
differential_scaler = DIFFERENTIAL_SCALER
|
||||
|
||||
unconditional_diff = (unconditional_latents - conditional_latents)
|
||||
unconditional_diff_noise = unconditional_diff * differential_scaler
|
||||
conditional_diff = (conditional_latents - unconditional_latents)
|
||||
conditional_diff_noise = conditional_diff * differential_scaler
|
||||
conditional_diff_noise = conditional_diff_noise.detach().requires_grad_(False)
|
||||
unconditional_diff_noise = unconditional_diff_noise.detach().requires_grad_(False)
|
||||
#
|
||||
baseline_conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
baseline_unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
conditional_noise = noise + unconditional_diff_noise
|
||||
unconditional_noise = noise + conditional_diff_noise
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
conditional_noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
unconditional_noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
# cat_baseline_noisy_latents = torch.cat(
|
||||
# [baseline_conditional_noisy_latents, baseline_unconditional_noisy_latents],
|
||||
# dim=0
|
||||
# )
|
||||
|
||||
# Disable the LoRA network so we can predict parent network knowledge without it
|
||||
# sd.network.is_active = False
|
||||
# sd.unet.eval()
|
||||
|
||||
# Predict noise to get a baseline of what the parent network wants to do with the latents + noise.
|
||||
# This acts as our control to preserve the unaltered parts of the image.
|
||||
# baseline_prediction = sd.predict_noise(
|
||||
# latents=cat_baseline_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
# conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
# timestep=cat_timesteps,
|
||||
# guidance_scale=1.0,
|
||||
# **pred_kwargs # adapter residuals in here
|
||||
# ).detach()
|
||||
|
||||
# conditional_baseline_prediction, unconditional_baseline_prediction = torch.chunk(baseline_prediction, 2, dim=0)
|
||||
|
||||
# negative_network_weights = [weight * -1.0 for weight in network_weight_list]
|
||||
# positive_network_weights = [weight * 1.0 for weight in network_weight_list]
|
||||
# cat_network_weight_list = positive_network_weights + negative_network_weights
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
# sd.network.is_active = True
|
||||
|
||||
# sd.network.multiplier = cat_network_weight_list
|
||||
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
# prediction = prediction - baseline_prediction
|
||||
|
||||
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
|
||||
# pred_pos = pred_pos - conditional_baseline_prediction
|
||||
# pred_neg = pred_neg - unconditional_baseline_prediction
|
||||
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
pred_pos.float(),
|
||||
conditional_noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
pred_loss = pred_loss.mean([1, 2, 3])
|
||||
|
||||
pred_neg_loss = torch.nn.functional.mse_loss(
|
||||
pred_neg.float(),
|
||||
unconditional_noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
pred_neg_loss = pred_neg_loss.mean([1, 2, 3])
|
||||
|
||||
loss = pred_loss + pred_neg_loss
|
||||
|
||||
loss = loss.mean()
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
def get_direct_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
unconditional_embeds: Optional[PromptEmbeds] = None,
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
# Perform targeted guidance (working title)
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
# target_noise,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
# sd.network.is_active = True
|
||||
|
||||
# sd.network.multiplier = network_weight_list
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
if unconditional_embeds is not None:
|
||||
unconditional_embeds = unconditional_embeds.to(device, dtype=dtype).detach()
|
||||
unconditional_embeds = concat_prompt_embeds([unconditional_embeds, unconditional_embeds])
|
||||
|
||||
prediction = sd.predict_noise(
|
||||
latents=torch.cat([unconditional_noisy_latents, conditional_noisy_latents]).to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=concat_prompt_embeds([conditional_embeds,conditional_embeds]).to(device, dtype=dtype).detach(),
|
||||
unconditional_embeddings=unconditional_embeds,
|
||||
timestep=torch.cat([timesteps, timesteps]),
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
noise_pred_uncond, noise_pred_cond = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
guidance_scale = 1.1
|
||||
guidance_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_cond - noise_pred_uncond
|
||||
)
|
||||
|
||||
guidance_loss = torch.nn.functional.mse_loss(
|
||||
guidance_pred.float(),
|
||||
noise.detach().float(),
|
||||
reduction="none"
|
||||
)
|
||||
if mask_multiplier is not None:
|
||||
guidance_loss = guidance_loss * mask_multiplier
|
||||
|
||||
guidance_loss = guidance_loss.mean([1, 2, 3])
|
||||
|
||||
guidance_loss = guidance_loss.mean()
|
||||
|
||||
# loss = guidance_loss + masked_noise_loss
|
||||
loss = guidance_loss
|
||||
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
# targeted
|
||||
def get_targeted_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
**kwargs
|
||||
):
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
# Encode the unconditional image into latents
|
||||
unconditional_noisy_latents = sd.noise_scheduler.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
)
|
||||
conditional_noisy_latents = sd.noise_scheduler.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
)
|
||||
|
||||
# was_network_active = self.network.is_active
|
||||
sd.network.is_active = False
|
||||
sd.unet.eval()
|
||||
|
||||
target_differential = unconditional_latents - conditional_latents
|
||||
# scale our loss by the differential scaler
|
||||
target_differential_abs = target_differential.abs()
|
||||
target_differential_abs_min = \
|
||||
target_differential_abs.min(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
target_differential_abs_max = \
|
||||
target_differential_abs.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
|
||||
min_guidance = 1.0
|
||||
max_guidance = 2.0
|
||||
|
||||
differential_scaler = value_map(
|
||||
target_differential_abs,
|
||||
target_differential_abs_min,
|
||||
target_differential_abs_max,
|
||||
min_guidance,
|
||||
max_guidance
|
||||
).detach()
|
||||
|
||||
|
||||
# With LoRA network bypassed, predict noise to get a baseline of what the network
|
||||
# wants to do with the latents + noise. Pass our target latents here for the input.
|
||||
target_unconditional = sd.predict_noise(
|
||||
latents=unconditional_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
).detach()
|
||||
prior_prediction_loss = torch.nn.functional.mse_loss(
|
||||
target_unconditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
).detach().clone()
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
sd.network.multiplier = network_weight_list + [x + -1.0 for x in network_weight_list]
|
||||
|
||||
# with LoRA active, predict the noise with the scaled differential latents added. This will allow us
|
||||
# the opportunity to predict the differential + noise that was added to the latents.
|
||||
prediction = sd.predict_noise(
|
||||
latents=torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0).to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=concat_prompt_embeds([conditional_embeds, conditional_embeds]).to(device, dtype=dtype).detach(),
|
||||
timestep=torch.cat([timesteps, timesteps], dim=0),
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
prediction_conditional, prediction_unconditional = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
conditional_loss = torch.nn.functional.mse_loss(
|
||||
prediction_conditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
unconditional_loss = torch.nn.functional.mse_loss(
|
||||
prediction_unconditional.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
positive_loss = torch.abs(
|
||||
conditional_loss.float() - prior_prediction_loss.float(),
|
||||
)
|
||||
# scale our loss by the differential scaler
|
||||
positive_loss = positive_loss * differential_scaler
|
||||
|
||||
positive_loss = positive_loss.mean([1, 2, 3])
|
||||
|
||||
polar_loss = torch.abs(
|
||||
conditional_loss.float() - unconditional_loss.float(),
|
||||
).mean([1, 2, 3])
|
||||
|
||||
|
||||
positive_loss = positive_loss.mean() + polar_loss.mean()
|
||||
|
||||
|
||||
positive_loss.backward()
|
||||
# loss = positive_loss.detach() + negative_loss.detach()
|
||||
loss = positive_loss.detach()
|
||||
|
||||
# add a grad so other backward does not fail
|
||||
loss.requires_grad_(True)
|
||||
|
||||
# restore network
|
||||
sd.network.multiplier = network_weight_list
|
||||
|
||||
return loss
|
||||
|
||||
def get_guided_loss_polarity(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
scaler=None,
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(dtype)
|
||||
noise = noise.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
target_pos = noise
|
||||
target_neg = noise
|
||||
|
||||
if sd.is_flow_matching:
|
||||
# set the timesteps for flow matching as linear since we will do weighing
|
||||
sd.noise_scheduler.set_train_timesteps(1000, device, linear=True)
|
||||
target_pos = (noise - conditional_latents).detach()
|
||||
target_neg = (noise - unconditional_latents).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
|
||||
negative_network_weights = [weight * -1.0 for weight in network_weight_list]
|
||||
positive_network_weights = [weight * 1.0 for weight in network_weight_list]
|
||||
cat_network_weight_list = positive_network_weights + negative_network_weights
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
|
||||
sd.network.multiplier = cat_network_weight_list
|
||||
|
||||
# do our prediction with LoRA active on the scaled guidance latents
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
|
||||
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
pred_pos.float(),
|
||||
target_pos.float(),
|
||||
reduction="none"
|
||||
)
|
||||
# pred_loss = pred_loss.mean([1, 2, 3])
|
||||
|
||||
pred_neg_loss = torch.nn.functional.mse_loss(
|
||||
pred_neg.float(),
|
||||
target_neg.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
loss = pred_loss + pred_neg_loss
|
||||
|
||||
if sd.is_flow_matching:
|
||||
timestep_weight = sd.noise_scheduler.get_weights_for_timesteps(timesteps).to(loss.device, dtype=loss.dtype).detach()
|
||||
loss = loss * timestep_weight
|
||||
|
||||
|
||||
loss = loss.mean([1, 2, 3])
|
||||
loss = loss.mean()
|
||||
if scaler is not None:
|
||||
scaler.scale(loss).backward()
|
||||
else:
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
|
||||
def get_guided_tnt(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
prior_pred: torch.Tensor = None,
|
||||
**kwargs
|
||||
):
|
||||
dtype = get_torch_dtype(sd.torch_dtype)
|
||||
device = sd.device_torch
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(dtype)
|
||||
noise = noise.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
|
||||
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
|
||||
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# double up everything to run it through all at once
|
||||
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
|
||||
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
|
||||
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
|
||||
|
||||
|
||||
# turn the LoRA network back on.
|
||||
sd.unet.train()
|
||||
if sd.network is not None:
|
||||
cat_network_weight_list = [weight for weight in network_weight_list * 2]
|
||||
sd.network.multiplier = cat_network_weight_list
|
||||
sd.network.is_active = True
|
||||
|
||||
|
||||
prediction = sd.predict_noise(
|
||||
latents=cat_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=cat_timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs # adapter residuals in here
|
||||
)
|
||||
this_prediction, that_prediction = torch.chunk(prediction, 2, dim=0)
|
||||
|
||||
this_loss = torch.nn.functional.mse_loss(
|
||||
this_prediction.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
that_loss = torch.nn.functional.mse_loss(
|
||||
that_prediction.float(),
|
||||
noise.float(),
|
||||
reduction="none"
|
||||
)
|
||||
|
||||
this_loss = this_loss.mean([1, 2, 3])
|
||||
# negative loss on that
|
||||
that_loss = -that_loss.mean([1, 2, 3])
|
||||
|
||||
with torch.no_grad():
|
||||
# match that loss with this loss so it is not a negative value and same scale
|
||||
that_loss_scaler = torch.abs(this_loss) / torch.abs(that_loss)
|
||||
|
||||
that_loss = that_loss * that_loss_scaler * 0.01
|
||||
|
||||
loss = this_loss + that_loss
|
||||
|
||||
loss = loss.mean()
|
||||
|
||||
loss.backward()
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
loss.requires_grad_(True)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
|
||||
# this processes all guidance losses based on the batch information
|
||||
def get_guidance_loss(
|
||||
noisy_latents: torch.Tensor,
|
||||
conditional_embeds: 'PromptEmbeds',
|
||||
match_adapter_assist: bool,
|
||||
network_weight_list: list,
|
||||
timesteps: torch.Tensor,
|
||||
pred_kwargs: dict,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
unconditional_embeds: Optional[PromptEmbeds] = None,
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
scaler=None,
|
||||
**kwargs
|
||||
):
|
||||
# TODO add others and process individual batch items separately
|
||||
guidance_type: GuidanceType = batch.file_items[0].dataset_config.guidance_type
|
||||
|
||||
if guidance_type == "targeted":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted guidance"
|
||||
return get_targeted_guidance_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "polarity":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
|
||||
return get_guided_loss_polarity(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
scaler=scaler,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "tnt":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
|
||||
return get_guided_tnt(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
prior_pred=prior_pred,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
elif guidance_type == "targeted_polarity":
|
||||
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted polarity guidance"
|
||||
return get_targeted_polarity_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "direct":
|
||||
return get_direct_guidance_loss(
|
||||
noisy_latents,
|
||||
conditional_embeds,
|
||||
match_adapter_assist,
|
||||
network_weight_list,
|
||||
timesteps,
|
||||
pred_kwargs,
|
||||
batch,
|
||||
noise,
|
||||
sd,
|
||||
unconditional_embeds=unconditional_embeds,
|
||||
mask_multiplier=mask_multiplier,
|
||||
prior_pred=prior_pred,
|
||||
**kwargs
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Guidance type {guidance_type} is not implemented")
|
||||
@@ -5,6 +5,7 @@ import json
|
||||
import os
|
||||
import io
|
||||
import struct
|
||||
import threading
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import cv2
|
||||
@@ -425,43 +426,63 @@ def main(argv=None):
|
||||
|
||||
|
||||
is_window_shown = False
|
||||
display_lock = threading.Lock()
|
||||
current_img = None
|
||||
update_event = threading.Event()
|
||||
|
||||
def update_image(img, name):
|
||||
global current_img
|
||||
with display_lock:
|
||||
current_img = (img, name)
|
||||
update_event.set()
|
||||
|
||||
def display_image_in_thread():
|
||||
global is_window_shown
|
||||
|
||||
def display_img():
|
||||
global current_img
|
||||
while True:
|
||||
update_event.wait()
|
||||
with display_lock:
|
||||
if current_img:
|
||||
img, name = current_img
|
||||
cv2.imshow(name, img)
|
||||
current_img = None
|
||||
update_event.clear()
|
||||
if cv2.waitKey(1) & 0xFF == 27: # Esc key to stop
|
||||
cv2.destroyAllWindows()
|
||||
print('\nESC pressed, stopping')
|
||||
break
|
||||
|
||||
if not is_window_shown:
|
||||
is_window_shown = True
|
||||
threading.Thread(target=display_img, daemon=True).start()
|
||||
|
||||
|
||||
def show_img(img, name='AI Toolkit'):
|
||||
global is_window_shown
|
||||
|
||||
img = np.clip(img, 0, 255).astype(np.uint8)
|
||||
cv2.imshow(name, img[:, :, ::-1])
|
||||
k = cv2.waitKey(10) & 0xFF
|
||||
if k == 27: # Esc key to stop
|
||||
print('\nESC pressed, stopping')
|
||||
raise KeyboardInterrupt
|
||||
update_image(img[:, :, ::-1], name)
|
||||
if not is_window_shown:
|
||||
is_window_shown = True
|
||||
|
||||
display_image_in_thread()
|
||||
|
||||
|
||||
def show_tensors(imgs: torch.Tensor, name='AI Toolkit'):
|
||||
# if rank is 4
|
||||
if len(imgs.shape) == 4:
|
||||
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
|
||||
else:
|
||||
img_list = [imgs]
|
||||
# put images side by side
|
||||
|
||||
img = torch.cat(img_list, dim=3)
|
||||
# img is -1 to 1, convert to 0 to 255
|
||||
img = img / 2 + 0.5
|
||||
img_numpy = img.to(torch.float32).detach().cpu().numpy()
|
||||
img_numpy = np.clip(img_numpy, 0, 1) * 255
|
||||
# convert to numpy Move channel to last
|
||||
img_numpy = img_numpy.transpose(0, 2, 3, 1)
|
||||
# convert to uint8
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
|
||||
show_img(img_numpy[0], name=name)
|
||||
|
||||
|
||||
def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit'):
|
||||
# decode latents
|
||||
if vae.device == 'cpu':
|
||||
vae.to(latents.device)
|
||||
latents = latents / vae.config['scaling_factor']
|
||||
@@ -469,12 +490,24 @@ def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit
|
||||
show_tensors(imgs, name=name)
|
||||
|
||||
|
||||
|
||||
def on_exit():
|
||||
if is_window_shown:
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
|
||||
def reduce_contrast(tensor, factor):
|
||||
# Ensure factor is between 0 and 1
|
||||
factor = max(0, min(factor, 1))
|
||||
|
||||
# Calculate the mean of the tensor
|
||||
mean = torch.mean(tensor)
|
||||
|
||||
# Reduce contrast
|
||||
adjusted_tensor = (tensor - mean) * factor + mean
|
||||
|
||||
# Clip values to ensure they stay within -1 to 1 range
|
||||
return torch.clamp(adjusted_tensor, -1.0, 1.0)
|
||||
|
||||
atexit.register(on_exit)
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
3039
toolkit/keymaps/stable_diffusion_vega.json
Normal file
3039
toolkit/keymaps/stable_diffusion_vega.json
Normal file
File diff suppressed because it is too large
Load Diff
BIN
toolkit/keymaps/stable_diffusion_vega_ldm_base.safetensors
Normal file
BIN
toolkit/keymaps/stable_diffusion_vega_ldm_base.safetensors
Normal file
Binary file not shown.
@@ -1,10 +1,13 @@
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import weakref
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from typing import List, Optional, Dict, Type, Union
|
||||
import torch
|
||||
from diffusers import UNet2DConditionModel, PixArtTransformer2DModel, AuraFlowTransformer2DModel
|
||||
from transformers import CLIPTextModel
|
||||
|
||||
from .config_modules import NetworkConfig
|
||||
@@ -15,6 +18,7 @@ from .paths import SD_SCRIPTS_ROOT
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from networks.lora import LoRANetwork, get_block_index
|
||||
from toolkit.models.DoRA import DoRAModule
|
||||
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
@@ -24,12 +28,14 @@ RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers
|
||||
# diffusers specific stuff
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear'
|
||||
'LoRACompatibleLinear',
|
||||
'QLinear',
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
'Conv2d',
|
||||
'LoRACompatibleConv'
|
||||
'LoRACompatibleConv',
|
||||
'QConv2d',
|
||||
]
|
||||
|
||||
class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
@@ -51,10 +57,12 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
use_bias: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
self.can_merge_in = True
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
ToolkitModuleMixin.__init__(self, network=network)
|
||||
torch.nn.Module.__init__(self)
|
||||
self.lora_name = lora_name
|
||||
self.orig_module_ref = weakref.ref(org_module)
|
||||
self.scalar = torch.tensor(1.0)
|
||||
# check if parent has bias. if not force use_bias to False
|
||||
if org_module.bias is None:
|
||||
@@ -111,10 +119,14 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
NUM_OF_BLOCKS = 12 # フルモデル相当でのup,downの層の数
|
||||
|
||||
UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
|
||||
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
|
||||
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
|
||||
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel", "ResnetBlock2D"]
|
||||
UNET_TARGET_REPLACE_MODULE = ["UNet2DConditionModel"]
|
||||
# UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
|
||||
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["UNet2DConditionModel"]
|
||||
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"]
|
||||
LORA_PREFIX_UNET = "lora_unet"
|
||||
PEFT_PREFIX_UNET = "unet"
|
||||
LORA_PREFIX_TEXT_ENCODER = "lora_te"
|
||||
|
||||
# SDXL: must starts with LORA_PREFIX_TEXT_ENCODER
|
||||
@@ -147,12 +159,23 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
train_unet: Optional[bool] = True,
|
||||
is_sdxl=False,
|
||||
is_v2=False,
|
||||
is_v3=False,
|
||||
is_pixart: bool = False,
|
||||
is_auraflow: bool = False,
|
||||
is_flux: bool = False,
|
||||
use_bias: bool = False,
|
||||
is_lorm: bool = False,
|
||||
ignore_if_contains = None,
|
||||
only_if_contains = None,
|
||||
parameter_threshold: float = 0.0,
|
||||
attn_only: bool = False,
|
||||
target_lin_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE,
|
||||
target_conv_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3,
|
||||
network_type: str = "lora",
|
||||
full_train_in_out: bool = False,
|
||||
transformer_only: bool = False,
|
||||
peft_format: bool = False,
|
||||
is_assistant_adapter: bool = False,
|
||||
**kwargs
|
||||
) -> None:
|
||||
"""
|
||||
@@ -177,6 +200,10 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if ignore_if_contains is None:
|
||||
ignore_if_contains = []
|
||||
self.ignore_if_contains = ignore_if_contains
|
||||
self.transformer_only = transformer_only
|
||||
|
||||
self.only_if_contains: Union[List, None] = only_if_contains
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
self.alpha = alpha
|
||||
self.conv_lora_dim = conv_lora_dim
|
||||
@@ -192,6 +219,30 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
self.multiplier = multiplier
|
||||
self.is_sdxl = is_sdxl
|
||||
self.is_v2 = is_v2
|
||||
self.is_v3 = is_v3
|
||||
self.is_pixart = is_pixart
|
||||
self.is_auraflow = is_auraflow
|
||||
self.is_flux = is_flux
|
||||
self.network_type = network_type
|
||||
self.is_assistant_adapter = is_assistant_adapter
|
||||
if self.network_type.lower() == "dora":
|
||||
self.module_class = DoRAModule
|
||||
module_class = DoRAModule
|
||||
|
||||
self.peft_format = peft_format
|
||||
|
||||
# always do peft for flux only for now
|
||||
if self.is_flux:
|
||||
self.peft_format = True
|
||||
|
||||
if self.peft_format:
|
||||
# no alpha for peft
|
||||
self.alpha = self.lora_dim
|
||||
alpha = self.alpha
|
||||
self.conv_alpha = self.conv_lora_dim
|
||||
conv_alpha = self.conv_alpha
|
||||
|
||||
self.full_train_in_out = full_train_in_out
|
||||
|
||||
if modules_dim is not None:
|
||||
print(f"create LoRA network from weights")
|
||||
@@ -219,8 +270,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
root_module: torch.nn.Module,
|
||||
target_replace_modules: List[torch.nn.Module],
|
||||
) -> List[LoRAModule]:
|
||||
unet_prefix = self.LORA_PREFIX_UNET
|
||||
if self.peft_format:
|
||||
unet_prefix = self.PEFT_PREFIX_UNET
|
||||
if is_pixart or is_v3 or is_auraflow or is_flux:
|
||||
unet_prefix = f"lora_transformer"
|
||||
if self.peft_format:
|
||||
unet_prefix = "transformer"
|
||||
|
||||
prefix = (
|
||||
self.LORA_PREFIX_UNET
|
||||
unet_prefix
|
||||
if is_unet
|
||||
else (
|
||||
self.LORA_PREFIX_TEXT_ENCODER
|
||||
@@ -230,6 +289,8 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
)
|
||||
loras = []
|
||||
skipped = []
|
||||
attached_modules = []
|
||||
lora_shape_dict = {}
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
for child_name, child_module in module.named_modules():
|
||||
@@ -237,6 +298,20 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
is_conv2d = child_module.__class__.__name__ in CONV_MODULES
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
|
||||
lora_name = [prefix, name, child_name]
|
||||
# filter out blank
|
||||
lora_name = [x for x in lora_name if x and x != ""]
|
||||
lora_name = ".".join(lora_name)
|
||||
# if it doesnt have a name, it wil have two dots
|
||||
lora_name.replace("..", ".")
|
||||
if self.peft_format:
|
||||
# we replace this on saving
|
||||
lora_name = lora_name.replace(".", "$$")
|
||||
else:
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
|
||||
skip = False
|
||||
if any([word in child_name for word in self.ignore_if_contains]):
|
||||
skip = True
|
||||
@@ -245,9 +320,17 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if count_parameters(child_module) < parameter_threshold:
|
||||
skip = True
|
||||
|
||||
if self.transformer_only and self.is_pixart and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_flux and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
if (is_linear or is_conv2d) and not skip:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
if self.only_if_contains is not None and not any([word in lora_name for word in self.only_if_contains]):
|
||||
continue
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
@@ -296,6 +379,8 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
use_bias=use_bias,
|
||||
)
|
||||
loras.append(lora)
|
||||
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape), list(lora.lora_up.weight.shape)
|
||||
]
|
||||
return loras, skipped
|
||||
|
||||
text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
|
||||
@@ -317,8 +402,12 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
index = None
|
||||
print(f"create LoRA for Text Encoder:")
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder,
|
||||
LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
|
||||
replace_modules = LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE
|
||||
|
||||
if self.is_pixart:
|
||||
replace_modules = ["T5EncoderModel"]
|
||||
|
||||
text_encoder_loras, skipped = create_modules(False, index, text_encoder, replace_modules)
|
||||
self.text_encoder_loras.extend(text_encoder_loras)
|
||||
skipped_te += skipped
|
||||
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
|
||||
@@ -328,6 +417,18 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if modules_dim is not None or self.conv_lora_dim is not None or conv_block_dims is not None:
|
||||
target_modules += target_conv_modules
|
||||
|
||||
if is_v3:
|
||||
target_modules = ["SD3Transformer2DModel"]
|
||||
|
||||
if is_pixart:
|
||||
target_modules = ["PixArtTransformer2DModel"]
|
||||
|
||||
if is_auraflow:
|
||||
target_modules = ["AuraFlowTransformer2DModel"]
|
||||
|
||||
if is_flux:
|
||||
target_modules = ["FluxTransformer2DModel"]
|
||||
|
||||
if train_unet:
|
||||
self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules)
|
||||
else:
|
||||
@@ -353,3 +454,49 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
for lora in self.text_encoder_loras + self.unet_loras:
|
||||
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
|
||||
names.add(lora.lora_name)
|
||||
|
||||
if self.full_train_in_out:
|
||||
print("full train in out")
|
||||
# we are going to retrain the main in out layers for VAE change usually
|
||||
if self.is_pixart:
|
||||
transformer: PixArtTransformer2DModel = unet
|
||||
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
|
||||
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
|
||||
|
||||
transformer.pos_embed = self.transformer_pos_embed
|
||||
transformer.proj_out = self.transformer_proj_out
|
||||
|
||||
elif self.is_auraflow:
|
||||
transformer: AuraFlowTransformer2DModel = unet
|
||||
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
|
||||
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
|
||||
|
||||
transformer.pos_embed = self.transformer_pos_embed
|
||||
transformer.proj_out = self.transformer_proj_out
|
||||
|
||||
else:
|
||||
unet: UNet2DConditionModel = unet
|
||||
unet_conv_in: torch.nn.Conv2d = unet.conv_in
|
||||
unet_conv_out: torch.nn.Conv2d = unet.conv_out
|
||||
|
||||
# clone these and replace their forwards with ours
|
||||
self.unet_conv_in = copy.deepcopy(unet_conv_in)
|
||||
self.unet_conv_out = copy.deepcopy(unet_conv_out)
|
||||
unet.conv_in = self.unet_conv_in
|
||||
unet.conv_out = self.unet_conv_out
|
||||
|
||||
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
|
||||
# call Lora prepare_optimizer_params
|
||||
all_params = super().prepare_optimizer_params(text_encoder_lr, unet_lr, default_lr)
|
||||
|
||||
if self.full_train_in_out:
|
||||
if self.is_pixart or self.is_auraflow or self.is_flux:
|
||||
all_params.append({"lr": unet_lr, "params": list(self.transformer_pos_embed.parameters())})
|
||||
all_params.append({"lr": unet_lr, "params": list(self.transformer_proj_out.parameters())})
|
||||
else:
|
||||
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_in.parameters())})
|
||||
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_out.parameters())})
|
||||
|
||||
return all_params
|
||||
|
||||
|
||||
|
||||
@@ -23,6 +23,8 @@ def get_meta_for_safetensors(meta: OrderedDict, name=None, add_software_info=Tru
|
||||
# if not float, int, bool, or str, convert to json string
|
||||
if not isinstance(value, str):
|
||||
save_meta[key] = json.dumps(value)
|
||||
# add the pt format
|
||||
save_meta["format"] = "pt"
|
||||
return save_meta
|
||||
|
||||
|
||||
|
||||
146
toolkit/models/DoRA.py
Normal file
146
toolkit/models/DoRA.py
Normal file
@@ -0,0 +1,146 @@
|
||||
#based off https://github.com/catid/dora/blob/main/dora.py
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from typing import TYPE_CHECKING, Union, List
|
||||
|
||||
from optimum.quanto import QBytesTensor, QTensor
|
||||
|
||||
from toolkit.network_mixins import ToolkitModuleMixin, ExtractableModuleMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
|
||||
# diffusers specific stuff
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear'
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
'Conv2d',
|
||||
'LoRACompatibleConv'
|
||||
]
|
||||
|
||||
def transpose(weight, fan_in_fan_out):
|
||||
if not fan_in_fan_out:
|
||||
return weight
|
||||
|
||||
if isinstance(weight, torch.nn.Parameter):
|
||||
return torch.nn.Parameter(weight.T)
|
||||
return weight.T
|
||||
|
||||
class DoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
# def __init__(self, d_in, d_out, rank=4, weight=None, bias=None):
|
||||
def __init__(
|
||||
self,
|
||||
lora_name,
|
||||
org_module: torch.nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
dropout=None,
|
||||
rank_dropout=None,
|
||||
module_dropout=None,
|
||||
network: 'LoRASpecialNetwork' = None,
|
||||
use_bias: bool = False,
|
||||
**kwargs
|
||||
):
|
||||
self.can_merge_in = False
|
||||
"""if alpha == 0 or None, alpha is rank (no scaling)."""
|
||||
ToolkitModuleMixin.__init__(self, network=network)
|
||||
torch.nn.Module.__init__(self)
|
||||
self.lora_name = lora_name
|
||||
self.scalar = torch.tensor(1.0)
|
||||
|
||||
self.lora_dim = lora_dim
|
||||
|
||||
if org_module.__class__.__name__ in CONV_MODULES:
|
||||
raise NotImplementedError("Convolutional layers are not supported yet")
|
||||
|
||||
if type(alpha) == torch.Tensor:
|
||||
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
|
||||
self.scale = alpha / self.lora_dim
|
||||
# self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える eng: treat as constant
|
||||
|
||||
self.multiplier: Union[float, List[float]] = multiplier
|
||||
# wrap the original module so it doesn't get weights updated
|
||||
self.org_module = [org_module]
|
||||
self.dropout = dropout
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
self.is_checkpointing = False
|
||||
|
||||
d_out = org_module.out_features
|
||||
d_in = org_module.in_features
|
||||
|
||||
std_dev = 1 / torch.sqrt(torch.tensor(self.lora_dim).float())
|
||||
# self.lora_up = nn.Parameter(torch.randn(d_out, self.lora_dim) * std_dev) # lora_A
|
||||
# self.lora_down = nn.Parameter(torch.zeros(self.lora_dim, d_in)) # lora_B
|
||||
self.lora_up = nn.Linear(self.lora_dim, d_out, bias=False) # lora_B
|
||||
# self.lora_up.weight.data = torch.randn_like(self.lora_up.weight.data) * std_dev
|
||||
self.lora_up.weight.data = torch.zeros_like(self.lora_up.weight.data)
|
||||
# self.lora_A[adapter_name] = nn.Linear(self.in_features, r, bias=False)
|
||||
# self.lora_B[adapter_name] = nn.Linear(r, self.out_features, bias=False)
|
||||
self.lora_down = nn.Linear(d_in, self.lora_dim, bias=False) # lora_A
|
||||
# self.lora_down.weight.data = torch.zeros_like(self.lora_down.weight.data)
|
||||
self.lora_down.weight.data = torch.randn_like(self.lora_down.weight.data) * std_dev
|
||||
|
||||
# m = Magnitude column-wise across output dimension
|
||||
weight = self.get_orig_weight()
|
||||
weight = weight.to(self.lora_up.weight.device, dtype=self.lora_up.weight.dtype)
|
||||
lora_weight = self.lora_up.weight @ self.lora_down.weight
|
||||
weight_norm = self._get_weight_norm(weight, lora_weight)
|
||||
self.magnitude = nn.Parameter(weight_norm.detach().clone(), requires_grad=True)
|
||||
|
||||
def apply_to(self):
|
||||
self.org_forward = self.org_module[0].forward
|
||||
self.org_module[0].forward = self.forward
|
||||
# del self.org_module
|
||||
|
||||
def get_orig_weight(self):
|
||||
weight = self.org_module[0].weight
|
||||
if isinstance(weight, QTensor) or isinstance(weight, QBytesTensor):
|
||||
return weight.dequantize().data.detach()
|
||||
else:
|
||||
return weight.data.detach()
|
||||
|
||||
def get_orig_bias(self):
|
||||
if hasattr(self.org_module[0], 'bias') and self.org_module[0].bias is not None:
|
||||
return self.org_module[0].bias.data.detach()
|
||||
return None
|
||||
|
||||
# def dora_forward(self, x, *args, **kwargs):
|
||||
# lora = torch.matmul(self.lora_A, self.lora_B)
|
||||
# adapted = self.get_orig_weight() + lora
|
||||
# column_norm = adapted.norm(p=2, dim=0, keepdim=True)
|
||||
# norm_adapted = adapted / column_norm
|
||||
# calc_weights = self.magnitude * norm_adapted
|
||||
# return F.linear(x, calc_weights, self.get_orig_bias())
|
||||
|
||||
def _get_weight_norm(self, weight, scaled_lora_weight) -> torch.Tensor:
|
||||
# calculate L2 norm of weight matrix, column-wise
|
||||
weight = weight + scaled_lora_weight.to(weight.device)
|
||||
weight_norm = torch.linalg.norm(weight, dim=1)
|
||||
return weight_norm
|
||||
|
||||
def apply_dora(self, x, scaled_lora_weight):
|
||||
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L192
|
||||
# lora weight is already scaled
|
||||
|
||||
# magnitude = self.lora_magnitude_vector[active_adapter]
|
||||
weight = self.get_orig_weight()
|
||||
weight = weight.to(scaled_lora_weight.device, dtype=scaled_lora_weight.dtype)
|
||||
weight_norm = self._get_weight_norm(weight, scaled_lora_weight)
|
||||
# see section 4.3 of DoRA (https://arxiv.org/abs/2402.09353)
|
||||
# "[...] we suggest treating ||V +∆V ||_c in
|
||||
# Eq. (5) as a constant, thereby detaching it from the gradient
|
||||
# graph. This means that while ||V + ∆V ||_c dynamically
|
||||
# reflects the updates of ∆V , it won’t receive any gradient
|
||||
# during backpropagation"
|
||||
weight_norm = weight_norm.detach()
|
||||
dora_weight = transpose(weight + scaled_lora_weight, False)
|
||||
return (self.magnitude / weight_norm - 1).view(1, -1) * F.linear(x.to(dora_weight.dtype), dora_weight)
|
||||
267
toolkit/models/LoRAFormer.py
Normal file
267
toolkit/models/LoRAFormer.py
Normal file
@@ -0,0 +1,267 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, d_model, nhead, dim_feedforward):
|
||||
super().__init__()
|
||||
self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
|
||||
self.cross_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
|
||||
self.feed_forward = nn.Sequential(
|
||||
nn.Linear(d_model, dim_feedforward),
|
||||
nn.ReLU(),
|
||||
nn.Linear(dim_feedforward, d_model)
|
||||
)
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.norm3 = nn.LayerNorm(d_model)
|
||||
|
||||
def forward(self, x, cross_attn_input):
|
||||
# Self-attention
|
||||
attn_output, _ = self.self_attn(x, x, x)
|
||||
x = self.norm1(x + attn_output)
|
||||
|
||||
# Cross-attention
|
||||
cross_attn_output, _ = self.cross_attn(x, cross_attn_input, cross_attn_input)
|
||||
x = self.norm2(x + cross_attn_output)
|
||||
|
||||
# Feed-forward
|
||||
ff_output = self.feed_forward(x)
|
||||
x = self.norm3(x + ff_output)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.down_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.up_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
# Initialize the network
|
||||
# num_blocks = 8
|
||||
# d_model = 1024 # Adjust as needed
|
||||
# nhead = 16 # Adjust as needed
|
||||
# dim_feedforward = 4096 # Adjust as needed
|
||||
# latent_dim = 1695744
|
||||
|
||||
class LoRAFormer(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_blocks,
|
||||
d_model=1024,
|
||||
nhead=16,
|
||||
dim_feedforward=4096,
|
||||
sd: 'StableDiffusion'=None,
|
||||
):
|
||||
super(LoRAFormer, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
self.latent = nn.Parameter(torch.randn(1, output_size))
|
||||
self.latent_proj = nn.Linear(output_size, d_model)
|
||||
self.blocks = nn.ModuleList([
|
||||
TransformerBlock(d_model, nhead, dim_feedforward)
|
||||
for _ in range(num_blocks)
|
||||
])
|
||||
self.final_proj = nn.Linear(d_model, output_size)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start+length])
|
||||
start += length
|
||||
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
}
|
||||
|
||||
127
toolkit/models/auraflow.py
Normal file
127
toolkit/models/auraflow.py
Normal file
@@ -0,0 +1,127 @@
|
||||
import math
|
||||
from functools import partial
|
||||
|
||||
from torch import nn
|
||||
import torch
|
||||
|
||||
|
||||
class AuraFlowPatchEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
height=224,
|
||||
width=224,
|
||||
patch_size=16,
|
||||
in_channels=3,
|
||||
embed_dim=768,
|
||||
pos_embed_max_size=None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.num_patches = (height // patch_size) * (width // patch_size)
|
||||
self.pos_embed_max_size = pos_embed_max_size
|
||||
|
||||
self.proj = nn.Linear(patch_size * patch_size * in_channels, embed_dim)
|
||||
self.pos_embed = nn.Parameter(torch.randn(1, pos_embed_max_size, embed_dim) * 0.1)
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.height, self.width = height // patch_size, width // patch_size
|
||||
self.base_size = height // patch_size
|
||||
|
||||
def forward(self, latent):
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
latent = latent.view(
|
||||
batch_size,
|
||||
num_channels,
|
||||
height // self.patch_size,
|
||||
self.patch_size,
|
||||
width // self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
latent = self.proj(latent)
|
||||
try:
|
||||
return latent + self.pos_embed
|
||||
except RuntimeError:
|
||||
raise RuntimeError(
|
||||
f"Positional embeddings are too small for the number of patches. "
|
||||
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
|
||||
)
|
||||
|
||||
|
||||
# comfy
|
||||
# def apply_pos_embeds(self, x, h, w):
|
||||
# h = (h + 1) // self.patch_size
|
||||
# w = (w + 1) // self.patch_size
|
||||
# max_dim = max(h, w)
|
||||
#
|
||||
# cur_dim = self.h_max
|
||||
# pos_encoding = self.positional_encoding.reshape(1, cur_dim, cur_dim, -1).to(device=x.device, dtype=x.dtype)
|
||||
#
|
||||
# if max_dim > cur_dim:
|
||||
# pos_encoding = F.interpolate(pos_encoding.movedim(-1, 1), (max_dim, max_dim), mode="bilinear").movedim(1,
|
||||
# -1)
|
||||
# cur_dim = max_dim
|
||||
#
|
||||
# from_h = (cur_dim - h) // 2
|
||||
# from_w = (cur_dim - w) // 2
|
||||
# pos_encoding = pos_encoding[:, from_h:from_h + h, from_w:from_w + w]
|
||||
# return x + pos_encoding.reshape(1, -1, self.positional_encoding.shape[-1])
|
||||
|
||||
# def patchify(self, x):
|
||||
# B, C, H, W = x.size()
|
||||
# pad_h = (self.patch_size - H % self.patch_size) % self.patch_size
|
||||
# pad_w = (self.patch_size - W % self.patch_size) % self.patch_size
|
||||
#
|
||||
# x = torch.nn.functional.pad(x, (0, pad_w, 0, pad_h), mode='reflect')
|
||||
# x = x.view(
|
||||
# B,
|
||||
# C,
|
||||
# (H + 1) // self.patch_size,
|
||||
# self.patch_size,
|
||||
# (W + 1) // self.patch_size,
|
||||
# self.patch_size,
|
||||
# )
|
||||
# x = x.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
# return x
|
||||
|
||||
def patch_auraflow_pos_embed(pos_embed):
|
||||
# we need to hijack the forward and replace with a custom one. Self is the model
|
||||
def new_forward(self, latent):
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
|
||||
# add padding to the latent to make it match pos_embed
|
||||
latent_size = height * width * num_channels / 16 # todo check where 16 comes from?
|
||||
pos_embed_size = self.pos_embed.shape[1]
|
||||
if latent_size < pos_embed_size:
|
||||
total_padding = int(pos_embed_size - math.floor(latent_size))
|
||||
total_padding = total_padding // 16
|
||||
pad_height = total_padding // 2
|
||||
pad_width = total_padding - pad_height
|
||||
# mirror padding on the right side
|
||||
padding = (0, pad_width, 0, pad_height)
|
||||
latent = torch.nn.functional.pad(latent, padding, mode='reflect')
|
||||
elif latent_size > pos_embed_size:
|
||||
amount_to_remove = latent_size - pos_embed_size
|
||||
latent = latent[:, :, :-amount_to_remove]
|
||||
|
||||
batch_size, num_channels, height, width = latent.size()
|
||||
|
||||
latent = latent.view(
|
||||
batch_size,
|
||||
num_channels,
|
||||
height // self.patch_size,
|
||||
self.patch_size,
|
||||
width // self.patch_size,
|
||||
self.patch_size,
|
||||
)
|
||||
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
|
||||
latent = self.proj(latent)
|
||||
try:
|
||||
return latent + self.pos_embed
|
||||
except RuntimeError:
|
||||
raise RuntimeError(
|
||||
f"Positional embeddings are too small for the number of patches. "
|
||||
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
|
||||
)
|
||||
|
||||
pos_embed.forward = partial(new_forward, pos_embed)
|
||||
162
toolkit/models/clip_fusion.py
Normal file
162
toolkit/models/clip_fusion.py
Normal file
@@ -0,0 +1,162 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from toolkit.models.zipper_resampler import ContextualAlphaMask
|
||||
|
||||
|
||||
# Conv1d MLP
|
||||
# MLP that can alternately be used as a conv1d on dim 1
|
||||
class MLPC(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim,
|
||||
out_dim,
|
||||
hidden_dim,
|
||||
do_conv=False,
|
||||
use_residual=True
|
||||
):
|
||||
super().__init__()
|
||||
self.do_conv = do_conv
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
# dont normalize if using conv
|
||||
if not do_conv:
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
|
||||
if do_conv:
|
||||
self.fc1 = nn.Conv1d(in_dim, hidden_dim, 1)
|
||||
self.fc2 = nn.Conv1d(hidden_dim, out_dim, 1)
|
||||
else:
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
if not self.do_conv:
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class ZipperBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_size,
|
||||
in_tokens,
|
||||
out_size,
|
||||
out_tokens,
|
||||
hidden_size,
|
||||
hidden_tokens,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_size = in_size
|
||||
self.in_tokens = in_tokens
|
||||
self.out_size = out_size
|
||||
self.out_tokens = out_tokens
|
||||
self.hidden_size = hidden_size
|
||||
self.hidden_tokens = hidden_tokens
|
||||
# permute to (batch_size, out_size, in_tokens)
|
||||
|
||||
self.zip_token = MLPC(
|
||||
in_dim=self.in_tokens,
|
||||
out_dim=self.out_tokens,
|
||||
hidden_dim=self.hidden_tokens,
|
||||
do_conv=True, # no need to permute
|
||||
use_residual=False
|
||||
)
|
||||
|
||||
# permute to (batch_size, out_tokens, out_size)
|
||||
|
||||
# in shpae: (batch_size, in_tokens, in_size)
|
||||
self.zip_size = MLPC(
|
||||
in_dim=self.in_size,
|
||||
out_dim=self.out_size,
|
||||
hidden_dim=self.hidden_size,
|
||||
use_residual=False
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.zip_token(x)
|
||||
x = self.zip_size(x)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# CLIPFusionModule
|
||||
# Fuses any size of vision and text embeddings into a single embedding.
|
||||
# remaps tokens and vectors.
|
||||
class CLIPFusionModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
text_hidden_size: int = 768,
|
||||
text_tokens: int = 77,
|
||||
vision_hidden_size: int = 1024,
|
||||
vision_tokens: int = 257,
|
||||
num_blocks: int = 1,
|
||||
):
|
||||
super(CLIPFusionModule, self).__init__()
|
||||
|
||||
self.text_hidden_size = text_hidden_size
|
||||
self.text_tokens = text_tokens
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
|
||||
self.resampler = ZipperBlock(
|
||||
in_size=self.vision_hidden_size,
|
||||
in_tokens=self.vision_tokens,
|
||||
out_size=self.text_hidden_size,
|
||||
out_tokens=self.text_tokens,
|
||||
hidden_size=self.vision_hidden_size * 2,
|
||||
hidden_tokens=self.vision_tokens * 2
|
||||
)
|
||||
|
||||
self.zipper_blocks = torch.nn.ModuleList([
|
||||
ZipperBlock(
|
||||
in_size=self.text_hidden_size * 2,
|
||||
in_tokens=self.text_tokens,
|
||||
out_size=self.text_hidden_size,
|
||||
out_tokens=self.text_tokens,
|
||||
hidden_size=self.text_hidden_size * 2,
|
||||
hidden_tokens=self.text_tokens * 2
|
||||
) for i in range(num_blocks)
|
||||
])
|
||||
|
||||
self.ctx_alpha = ContextualAlphaMask(
|
||||
dim=self.text_hidden_size,
|
||||
)
|
||||
|
||||
self.alpha = nn.Parameter(torch.zeros([text_tokens]) + 0.01)
|
||||
|
||||
def forward(self, text_embeds, vision_embeds):
|
||||
# text_embeds = (batch_size, 77, 768)
|
||||
# vision_embeds = (batch_size, 257, 1024)
|
||||
# output = (batch_size, 77, 768)
|
||||
|
||||
vision_embeds = self.resampler(vision_embeds)
|
||||
x = vision_embeds
|
||||
for i, block in enumerate(self.zipper_blocks):
|
||||
res = x
|
||||
x = torch.cat([text_embeds, x], dim=-1)
|
||||
x = block(x)
|
||||
x = x + res
|
||||
|
||||
# alpha mask
|
||||
ctx_alpha = self.ctx_alpha(text_embeds)
|
||||
# reshape alpha to (1, 77, 1)
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
|
||||
|
||||
x = ctx_alpha * x * alpha
|
||||
|
||||
x = x + text_embeds
|
||||
|
||||
return x
|
||||
123
toolkit/models/clip_pre_processor.py
Normal file
123
toolkit/models/clip_pre_processor.py
Normal file
@@ -0,0 +1,123 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class UpsampleBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.conv_in = nn.Sequential(
|
||||
nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1),
|
||||
nn.GELU()
|
||||
)
|
||||
self.conv_up = nn.Sequential(
|
||||
nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2),
|
||||
nn.GELU()
|
||||
)
|
||||
|
||||
self.conv_out = nn.Sequential(
|
||||
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
x = self.conv_up(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class CLIPImagePreProcessor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size=896,
|
||||
clip_input_size=224,
|
||||
downscale_factor: int = 16,
|
||||
):
|
||||
super().__init__()
|
||||
# make sure they are evenly divisible
|
||||
assert input_size % clip_input_size == 0
|
||||
in_channels = 3
|
||||
|
||||
self.input_size = input_size
|
||||
self.clip_input_size = clip_input_size
|
||||
self.downscale_factor = downscale_factor
|
||||
|
||||
subpixel_channels = in_channels * downscale_factor ** 2 # 3 * 16 ** 2 = 768
|
||||
channels = subpixel_channels
|
||||
|
||||
upscale_factor = downscale_factor / int((input_size / clip_input_size)) # 16 / (896 / 224) = 4
|
||||
|
||||
num_upsample_blocks = int(upscale_factor // 2) # 4 // 2 = 2
|
||||
|
||||
# make the residual down up blocks
|
||||
self.upsample_blocks = nn.ModuleList()
|
||||
self.subpixel_blocks = nn.ModuleList()
|
||||
current_channels = channels
|
||||
current_downscale = downscale_factor
|
||||
for _ in range(num_upsample_blocks):
|
||||
# determine the reshuffled channel count for this dimension
|
||||
output_downscale = current_downscale // 2
|
||||
out_channels = in_channels * output_downscale ** 2
|
||||
# out_channels = current_channels // 2
|
||||
self.upsample_blocks.append(UpsampleBlock(current_channels, out_channels))
|
||||
current_channels = out_channels
|
||||
current_downscale = output_downscale
|
||||
self.subpixel_blocks.append(nn.PixelUnshuffle(current_downscale))
|
||||
|
||||
# (bs, 768, 56, 56) -> (bs, 192, 112, 112)
|
||||
# (bs, 192, 112, 112) -> (bs, 48, 224, 224)
|
||||
|
||||
self.conv_out = nn.Conv2d(
|
||||
current_channels,
|
||||
out_channels=3,
|
||||
kernel_size=3,
|
||||
padding=1
|
||||
) # (bs, 48, 224, 224) -> (bs, 3, 224, 224)
|
||||
|
||||
# do a pooling layer to downscale the input to 1/3 of the size
|
||||
# (bs, 3, 896, 896) -> (bs, 3, 224, 224)
|
||||
kernel_size = input_size // clip_input_size
|
||||
self.res_down = nn.AvgPool2d(
|
||||
kernel_size=kernel_size,
|
||||
stride=kernel_size
|
||||
) # (bs, 3, 896, 896) -> (bs, 3, 224, 224)
|
||||
|
||||
# make a blending for output residual with near 0 weight
|
||||
self.res_blend = nn.Parameter(torch.tensor(0.001)) # (bs, 3, 224, 224) -> (bs, 3, 224, 224)
|
||||
|
||||
self.unshuffle = nn.PixelUnshuffle(downscale_factor) # (bs, 3, 896, 896) -> (bs, 768, 56, 56)
|
||||
|
||||
self.conv_in = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
subpixel_channels,
|
||||
channels,
|
||||
kernel_size=3,
|
||||
padding=1
|
||||
),
|
||||
nn.GELU()
|
||||
) # (bs, 768, 56, 56) -> (bs, 768, 56, 56)
|
||||
|
||||
# make 2 deep blocks
|
||||
|
||||
def forward(self, x):
|
||||
inputs = x
|
||||
# resize to input_size x input_size
|
||||
x = nn.functional.interpolate(x, size=(self.input_size, self.input_size), mode='bicubic')
|
||||
|
||||
res = self.res_down(inputs)
|
||||
|
||||
x = self.unshuffle(x)
|
||||
x = self.conv_in(x)
|
||||
for up, subpixel in zip(self.upsample_blocks, self.subpixel_blocks):
|
||||
x = up(x)
|
||||
block_res = subpixel(inputs)
|
||||
x = x + block_res
|
||||
x = self.conv_out(x)
|
||||
# blend residual
|
||||
x = x * self.res_blend + res
|
||||
return x
|
||||
359
toolkit/models/ilora.py
Normal file
359
toolkit/models/ilora.py
Normal file
@@ -0,0 +1,359 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
class LoRAGenerator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 768, # projection dimension
|
||||
hidden_size: int = 768,
|
||||
head_size: int = 512,
|
||||
num_heads: int = 1,
|
||||
num_mlp_layers: int = 1,
|
||||
output_size: int = 768,
|
||||
dropout: float = 0.0
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.num_heads = num_heads
|
||||
self.simple = False
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
if self.simple:
|
||||
self.head = nn.Linear(input_size, head_size, bias=False)
|
||||
else:
|
||||
self.lin_in = nn.Linear(input_size, hidden_size)
|
||||
|
||||
self.mlp_blocks = nn.Sequential(*[
|
||||
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in range(num_mlp_layers)
|
||||
])
|
||||
self.head = nn.Linear(hidden_size, head_size, bias=False)
|
||||
self.norm = nn.LayerNorm(head_size)
|
||||
|
||||
if num_heads == 1:
|
||||
self.output = nn.Linear(head_size, self.output_size)
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
self.output.weight.data *= 0.01
|
||||
else:
|
||||
head_output_size = output_size // num_heads
|
||||
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
for output in self.outputs:
|
||||
output.weight.data *= 0.01
|
||||
|
||||
# allow get device
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def forward(self, embedding):
|
||||
if len(embedding.shape) == 2:
|
||||
embedding = embedding.unsqueeze(1)
|
||||
|
||||
x = embedding
|
||||
|
||||
if not self.simple:
|
||||
x = self.lin_in(embedding)
|
||||
x = self.mlp_blocks(x)
|
||||
x = self.head(x)
|
||||
x = self.norm(x)
|
||||
|
||||
if self.num_heads == 1:
|
||||
x = self.output(x)
|
||||
else:
|
||||
out_chunks = torch.chunk(x, self.num_heads, dim=1)
|
||||
x = []
|
||||
for out_layer, chunk in zip(self.outputs, out_chunks):
|
||||
x.append(out_layer(chunk))
|
||||
x = torch.cat(x, dim=-1)
|
||||
|
||||
return x.squeeze(1)
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.down_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
org_module = self.lora_module_ref().orig_module_ref()
|
||||
stride = org_module.stride
|
||||
padding = org_module.padding
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding, stride=stride)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
weight_chunk = weight_chunk.view(self.up_shape)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class InstantLoRAModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vision_hidden_size: int,
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion'
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
number_formatted_output_size = "{:,}".format(output_size)
|
||||
|
||||
print(f" ILORA output size: {number_formatted_output_size}")
|
||||
|
||||
# if not evenly divisible, error
|
||||
if self.output_size % self.num_heads != 0:
|
||||
raise ValueError("Output size must be divisible by the number of heads")
|
||||
|
||||
self.head_output_size = self.output_size // self.num_heads
|
||||
|
||||
if vision_tokens > 1:
|
||||
self.resampler = Resampler(
|
||||
dim=vision_hidden_size,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=12,
|
||||
num_queries=num_heads, # output tokens
|
||||
embedding_dim=vision_hidden_size,
|
||||
max_seq_len=vision_tokens,
|
||||
output_dim=head_dim,
|
||||
apply_pos_emb=True, # this is new
|
||||
ff_mult=4
|
||||
)
|
||||
|
||||
self.proj_module = LoRAGenerator(
|
||||
input_size=head_dim,
|
||||
hidden_size=head_dim,
|
||||
head_size=head_dim,
|
||||
num_mlp_layers=1,
|
||||
num_heads=self.num_heads,
|
||||
output_size=self.output_size,
|
||||
)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start+length])
|
||||
start += length
|
||||
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
}
|
||||
|
||||
419
toolkit/models/ilora2.py
Normal file
419
toolkit/models/ilora2.py
Normal file
@@ -0,0 +1,419 @@
|
||||
import math
|
||||
import weakref
|
||||
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import TYPE_CHECKING, List, Dict, Any
|
||||
from toolkit.models.clip_fusion import ZipperBlock
|
||||
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
|
||||
import sys
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from collections import OrderedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class LoRAGenerator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_size: int = 768, # projection dimension
|
||||
hidden_size: int = 768,
|
||||
head_size: int = 512,
|
||||
num_heads: int = 1,
|
||||
num_mlp_layers: int = 1,
|
||||
output_size: int = 768,
|
||||
dropout: float = 0.0
|
||||
):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.num_heads = num_heads
|
||||
self.simple = False
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
if self.simple:
|
||||
self.head = nn.Linear(input_size, head_size, bias=False)
|
||||
else:
|
||||
self.lin_in = nn.Linear(input_size, hidden_size)
|
||||
|
||||
self.mlp_blocks = nn.Sequential(*[
|
||||
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in
|
||||
range(num_mlp_layers)
|
||||
])
|
||||
self.head = nn.Linear(hidden_size, head_size, bias=False)
|
||||
self.norm = nn.LayerNorm(head_size)
|
||||
|
||||
if num_heads == 1:
|
||||
self.output = nn.Linear(head_size, self.output_size)
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
self.output.weight.data *= 0.01
|
||||
else:
|
||||
head_output_size = output_size // num_heads
|
||||
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
|
||||
# for each output block. multiply weights by 0.01
|
||||
with torch.no_grad():
|
||||
for output in self.outputs:
|
||||
output.weight.data *= 0.01
|
||||
|
||||
# allow get device
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
def forward(self, embedding):
|
||||
if len(embedding.shape) == 2:
|
||||
embedding = embedding.unsqueeze(1)
|
||||
|
||||
x = embedding
|
||||
|
||||
if not self.simple:
|
||||
x = self.lin_in(embedding)
|
||||
x = self.mlp_blocks(x)
|
||||
x = self.head(x)
|
||||
x = self.norm(x)
|
||||
|
||||
if self.num_heads == 1:
|
||||
x = self.output(x)
|
||||
else:
|
||||
out_chunks = torch.chunk(x, self.num_heads, dim=1)
|
||||
x = []
|
||||
for out_layer, chunk in zip(self.outputs, out_chunks):
|
||||
x.append(out_layer(chunk))
|
||||
x = torch.cat(x, dim=-1)
|
||||
|
||||
return x.squeeze(1)
|
||||
|
||||
|
||||
class InstantLoRAMidModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
index: int,
|
||||
lora_module: 'LoRAModule',
|
||||
instant_lora_module: 'InstantLoRAModule',
|
||||
up_shape: list = None,
|
||||
down_shape: list = None,
|
||||
):
|
||||
super(InstantLoRAMidModule, self).__init__()
|
||||
self.up_shape = up_shape
|
||||
self.down_shape = down_shape
|
||||
self.index = index
|
||||
self.lora_module_ref = weakref.ref(lora_module)
|
||||
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
|
||||
|
||||
self.do_up = instant_lora_module.config.ilora_up
|
||||
self.do_down = instant_lora_module.config.ilora_down
|
||||
self.do_mid = instant_lora_module.config.ilora_mid
|
||||
|
||||
self.down_dim = self.down_shape[1] if self.do_down else 0
|
||||
self.mid_dim = self.up_shape[1] if self.do_mid else 0
|
||||
self.out_dim = self.up_shape[0] if self.do_up else 0
|
||||
|
||||
self.embed = None
|
||||
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
if not self.do_down:
|
||||
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
down_weight = self.embed[:, :self.down_dim]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if down_weight.shape[0] * 2 == batch_size:
|
||||
down_weight = torch.cat([down_weight] * 2, dim=0)
|
||||
|
||||
try:
|
||||
if len(x.shape) == 4:
|
||||
# conv
|
||||
down_weight = down_weight.view(batch_size, -1, 1, 1)
|
||||
if x.shape[1] != down_weight.shape[1]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
elif len(x.shape) == 2:
|
||||
down_weight = down_weight.view(batch_size, -1)
|
||||
if x.shape[1] != down_weight.shape[1]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
else:
|
||||
down_weight = down_weight.view(batch_size, 1, -1)
|
||||
if x.shape[2] != down_weight.shape[2]:
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
x = x * down_weight
|
||||
x = self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
|
||||
|
||||
return x
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
# do mid here
|
||||
x = self.mid_forward(x, *args, **kwargs)
|
||||
if not self.do_up:
|
||||
return self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
up_weight = self.embed[:, -self.out_dim:]
|
||||
|
||||
batch_size = x.shape[0]
|
||||
|
||||
# unconditional
|
||||
if up_weight.shape[0] * 2 == batch_size:
|
||||
up_weight = torch.cat([up_weight] * 2, dim=0)
|
||||
|
||||
try:
|
||||
if len(x.shape) == 4:
|
||||
# conv
|
||||
up_weight = up_weight.view(batch_size, -1, 1, 1)
|
||||
elif len(x.shape) == 2:
|
||||
up_weight = up_weight.view(batch_size, -1)
|
||||
else:
|
||||
up_weight = up_weight.view(batch_size, 1, -1)
|
||||
x = self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
|
||||
x = x * up_weight
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise ValueError(f"Up weight shape not understood: {up_weight.shape} {x.shape}")
|
||||
|
||||
return x
|
||||
|
||||
def mid_forward(self, x, *args, **kwargs):
|
||||
if not self.do_mid:
|
||||
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
|
||||
batch_size = x.shape[0]
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
mid_weight = self.embed[:, self.down_dim:self.down_dim + self.mid_dim * self.mid_dim]
|
||||
|
||||
# unconditional
|
||||
if mid_weight.shape[0] * 2 == batch_size:
|
||||
mid_weight = torch.cat([mid_weight] * 2, dim=0)
|
||||
|
||||
weight_chunks = torch.chunk(mid_weight, batch_size, dim=0)
|
||||
x_chunks = torch.chunk(x, batch_size, dim=0)
|
||||
|
||||
x_out = []
|
||||
for i in range(batch_size):
|
||||
weight_chunk = weight_chunks[i]
|
||||
x_chunk = x_chunks[i]
|
||||
# reshape
|
||||
if len(x_chunk.shape) == 4:
|
||||
# conv
|
||||
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim, 1, 1)
|
||||
else:
|
||||
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim)
|
||||
# check if is conv or linear
|
||||
if len(weight_chunk.shape) == 4:
|
||||
padding = 0
|
||||
if weight_chunk.shape[-1] == 3:
|
||||
padding = 1
|
||||
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
|
||||
else:
|
||||
# run a simple linear layer with the down weight
|
||||
x_chunk = x_chunk @ weight_chunk.T
|
||||
x_out.append(x_chunk)
|
||||
x = torch.cat(x_out, dim=0)
|
||||
return x
|
||||
|
||||
|
||||
class InstantLoRAModule(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
vision_hidden_size: int,
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion',
|
||||
config: AdapterConfig
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.dim = sd.network.lora_dim
|
||||
self.vision_hidden_size = vision_hidden_size
|
||||
self.vision_tokens = vision_tokens
|
||||
self.head_dim = head_dim
|
||||
self.num_heads = num_heads
|
||||
|
||||
self.config: AdapterConfig = config
|
||||
|
||||
# stores the projection vector. Grabbed by modules
|
||||
self.img_embeds: List[torch.Tensor] = None
|
||||
|
||||
# disable merging in. It is slower on inference
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
self.ilora_modules = torch.nn.ModuleList()
|
||||
|
||||
lora_modules = self.sd_ref().network.get_all_modules()
|
||||
|
||||
output_size = 0
|
||||
|
||||
self.embed_lengths = []
|
||||
self.weight_mapping = []
|
||||
|
||||
for idx, lora_module in enumerate(lora_modules):
|
||||
module_dict = lora_module.state_dict()
|
||||
down_shape = list(module_dict['lora_down.weight'].shape)
|
||||
up_shape = list(module_dict['lora_up.weight'].shape)
|
||||
|
||||
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
|
||||
|
||||
#
|
||||
# module_size = math.prod(down_shape) + math.prod(up_shape)
|
||||
|
||||
# conv weight shape is (out_channels, in_channels, kernel_size, kernel_size)
|
||||
# linear weight shape is (out_features, in_features)
|
||||
|
||||
# just doing in dim and out dim
|
||||
in_dim = down_shape[1] if self.config.ilora_down else 0
|
||||
mid_dim = down_shape[0] * down_shape[0] if self.config.ilora_mid else 0
|
||||
out_dim = up_shape[0] if self.config.ilora_up else 0
|
||||
module_size = in_dim + mid_dim + out_dim
|
||||
|
||||
output_size += module_size
|
||||
self.embed_lengths.append(module_size)
|
||||
|
||||
# add a new mid module that will take the original forward and add a vector to it
|
||||
# this will be used to add the vector to the original forward
|
||||
instant_module = InstantLoRAMidModule(
|
||||
idx,
|
||||
lora_module,
|
||||
self,
|
||||
up_shape=up_shape,
|
||||
down_shape=down_shape
|
||||
)
|
||||
|
||||
self.ilora_modules.append(instant_module)
|
||||
|
||||
# replace the LoRA forwards
|
||||
lora_module.lora_down.orig_forward = lora_module.lora_down.forward
|
||||
lora_module.lora_down.forward = instant_module.down_forward
|
||||
lora_module.lora_up.orig_forward = lora_module.lora_up.forward
|
||||
lora_module.lora_up.forward = instant_module.up_forward
|
||||
|
||||
self.output_size = output_size
|
||||
|
||||
number_formatted_output_size = "{:,}".format(output_size)
|
||||
|
||||
print(f" ILORA output size: {number_formatted_output_size}")
|
||||
|
||||
# if not evenly divisible, error
|
||||
if self.output_size % self.num_heads != 0:
|
||||
raise ValueError("Output size must be divisible by the number of heads")
|
||||
|
||||
self.head_output_size = self.output_size // self.num_heads
|
||||
|
||||
if vision_tokens > 1:
|
||||
self.resampler = Resampler(
|
||||
dim=vision_hidden_size,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=12,
|
||||
num_queries=num_heads, # output tokens
|
||||
embedding_dim=vision_hidden_size,
|
||||
max_seq_len=vision_tokens,
|
||||
output_dim=head_dim,
|
||||
apply_pos_emb=True, # this is new
|
||||
ff_mult=4
|
||||
)
|
||||
|
||||
self.proj_module = LoRAGenerator(
|
||||
input_size=head_dim,
|
||||
hidden_size=head_dim,
|
||||
head_size=head_dim,
|
||||
num_mlp_layers=1,
|
||||
num_heads=self.num_heads,
|
||||
output_size=self.output_size,
|
||||
)
|
||||
|
||||
self.migrate_weight_mapping()
|
||||
|
||||
def migrate_weight_mapping(self):
|
||||
return
|
||||
# # changes the names of the modules to common ones
|
||||
# keymap = self.sd_ref().network.get_keymap()
|
||||
# save_keymap = {}
|
||||
# if keymap is not None:
|
||||
# for ldm_key, diffusers_key in keymap.items():
|
||||
# # invert them
|
||||
# save_keymap[diffusers_key] = ldm_key
|
||||
#
|
||||
# new_keymap = {}
|
||||
# for key, value in self.weight_mapping:
|
||||
# if key in save_keymap:
|
||||
# new_keymap[save_keymap[key]] = value
|
||||
# else:
|
||||
# print(f"Key {key} not found in keymap")
|
||||
# new_keymap[key] = value
|
||||
# self.weight_mapping = new_keymap
|
||||
# else:
|
||||
# print("No keymap found. Using default names")
|
||||
# return
|
||||
|
||||
def forward(self, img_embeds):
|
||||
# expand token rank if only rank 2
|
||||
if len(img_embeds.shape) == 2:
|
||||
img_embeds = img_embeds.unsqueeze(1)
|
||||
|
||||
# resample the image embeddings
|
||||
img_embeds = self.resampler(img_embeds)
|
||||
img_embeds = self.proj_module(img_embeds)
|
||||
if len(img_embeds.shape) == 3:
|
||||
# merge the heads
|
||||
img_embeds = img_embeds.mean(dim=1)
|
||||
|
||||
self.img_embeds = []
|
||||
# get all the slices
|
||||
start = 0
|
||||
for length in self.embed_lengths:
|
||||
self.img_embeds.append(img_embeds[:, start:start + length])
|
||||
start += length
|
||||
|
||||
def get_additional_save_metadata(self) -> Dict[str, Any]:
|
||||
# save the weight mapping
|
||||
return {
|
||||
"weight_mapping": self.weight_mapping,
|
||||
"num_heads": self.num_heads,
|
||||
"vision_hidden_size": self.vision_hidden_size,
|
||||
"head_dim": self.head_dim,
|
||||
"vision_tokens": self.vision_tokens,
|
||||
"output_size": self.output_size,
|
||||
"do_up": self.config.ilora_up,
|
||||
"do_mid": self.config.ilora_mid,
|
||||
"do_down": self.config.ilora_down,
|
||||
}
|
||||
402
toolkit/models/single_value_adapter.py
Normal file
402
toolkit/models/single_value_adapter.py
Normal file
@@ -0,0 +1,402 @@
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Union, TYPE_CHECKING
|
||||
|
||||
from diffusers import Transformer2DModel
|
||||
from transformers import T5EncoderModel, CLIPTextModel, CLIPTokenizer, T5Tokenizer, CLIPVisionModelWithProjection
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
class AttnProcessor2_0(torch.nn.Module):
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size=None,
|
||||
cross_attention_dim=None,
|
||||
):
|
||||
super().__init__()
|
||||
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.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
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)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
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 SingleValueAdapterAttnProcessor(nn.Module):
|
||||
r"""
|
||||
Attention processor for Custom TE for PyTorch 2.0.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
scale (`float`, defaults to 1.0):
|
||||
the weight scale of image prompt.
|
||||
adapter
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, adapter=None,
|
||||
adapter_hidden_size=None, has_bias=False, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
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.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.adapter_hidden_size = adapter_hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
|
||||
self.to_k_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
self.to_v_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
# return False
|
||||
|
||||
@property
|
||||
def unconditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().unconditional_embeds
|
||||
|
||||
@property
|
||||
def conditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().conditional_embeds
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
is_active = self.adapter_ref().is_active
|
||||
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)
|
||||
|
||||
# will be none if disabled
|
||||
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
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# only use one TE or the other. If our adapter is active only use ours
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
|
||||
adapter_hidden_states = self.conditional_embeds
|
||||
if adapter_hidden_states.shape[0] < batch_size:
|
||||
# doing cfg
|
||||
adapter_hidden_states = torch.cat([
|
||||
self.unconditional_embeds,
|
||||
adapter_hidden_states
|
||||
], dim=0)
|
||||
# needs to be shape (batch, 1, 1)
|
||||
if len(adapter_hidden_states.shape) == 2:
|
||||
adapter_hidden_states = adapter_hidden_states.unsqueeze(1)
|
||||
# conditional_batch_size = adapter_hidden_states.shape[0]
|
||||
# conditional_query = query
|
||||
|
||||
# for ip-adapter
|
||||
vd_key = self.to_k_adapter(adapter_hidden_states)
|
||||
vd_value = self.to_v_adapter(adapter_hidden_states)
|
||||
|
||||
vd_key = vd_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
vd_value = vd_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
|
||||
vd_hidden_states = F.scaled_dot_product_attention(
|
||||
query, vd_key, vd_value, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
vd_hidden_states = vd_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
vd_hidden_states = vd_hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states = hidden_states + self.scale * vd_hidden_states
|
||||
|
||||
|
||||
# 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 SingleValueAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
num_values: int = 1,
|
||||
):
|
||||
super(SingleValueAdapter, self).__init__()
|
||||
is_pixart = sd.is_pixart
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.token_size = num_values
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
|
||||
attn_processor_keys = []
|
||||
if is_pixart:
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
|
||||
|
||||
# cross attention
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
|
||||
|
||||
else:
|
||||
attn_processor_keys = list(sd.unet.attn_processors.keys())
|
||||
|
||||
for name in attn_processor_keys:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") else sd.unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = sd.unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(sd.unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
hidden_size = sd.unet.config['cross_attention_dim']
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AttnProcessor2_0()
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
|
||||
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
|
||||
# if is_pixart:
|
||||
# to_k_bias = unet_sd[layer_name + ".to_k.bias"]
|
||||
# to_v_bias = unet_sd[layer_name + ".to_v.bias"]
|
||||
# else:
|
||||
# to_k_bias = None
|
||||
# to_v_bias = None
|
||||
|
||||
# add zero padding to the adapter
|
||||
if to_k_adapter.shape[1] < self.token_size:
|
||||
to_k_adapter = torch.cat([
|
||||
to_k_adapter,
|
||||
torch.randn(to_k_adapter.shape[0], self.token_size - to_k_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
to_v_adapter = torch.cat([
|
||||
to_v_adapter,
|
||||
torch.randn(to_v_adapter.shape[0], self.token_size - to_v_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
# if is_pixart:
|
||||
# to_k_bias = torch.cat([
|
||||
# to_k_bias,
|
||||
# torch.zeros(self.token_size - to_k_adapter.shape[1]).to(
|
||||
# to_k_adapter.device, dtype=to_k_adapter.dtype)
|
||||
# ],
|
||||
# dim=0
|
||||
# )
|
||||
# to_v_bias = torch.cat([
|
||||
# to_v_bias,
|
||||
# torch.zeros(self.token_size - to_v_adapter.shape[1]).to(
|
||||
# to_k_adapter.device, dtype=to_k_adapter.dtype)
|
||||
# ],
|
||||
# dim=0
|
||||
# )
|
||||
elif to_k_adapter.shape[1] > self.token_size:
|
||||
to_k_adapter = to_k_adapter[:, :self.token_size]
|
||||
to_v_adapter = to_v_adapter[:, :self.token_size]
|
||||
# if is_pixart:
|
||||
# to_k_bias = to_k_bias[:self.token_size]
|
||||
# to_v_bias = to_v_bias[:self.token_size]
|
||||
else:
|
||||
to_k_adapter = to_k_adapter
|
||||
to_v_adapter = to_v_adapter
|
||||
# if is_pixart:
|
||||
# to_k_bias = to_k_bias
|
||||
# to_v_bias = to_v_bias
|
||||
|
||||
weights = {
|
||||
"to_k_adapter.weight": to_k_adapter * 0.01,
|
||||
"to_v_adapter.weight": to_v_adapter * 0.01,
|
||||
}
|
||||
# if is_pixart:
|
||||
# weights["to_k_adapter.bias"] = to_k_bias
|
||||
# weights["to_v_adapter.bias"] = to_v_bias
|
||||
|
||||
attn_procs[name] = SingleValueAdapterAttnProcessor(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
has_bias=False,
|
||||
)
|
||||
attn_procs[name].load_state_dict(weights)
|
||||
if self.sd_ref().is_pixart:
|
||||
# we have to set them ourselves
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn1.processor = attn_procs[f"transformer_blocks.{i}.attn1"]
|
||||
module.attn2.processor = attn_procs[f"transformer_blocks.{i}.attn2"]
|
||||
self.adapter_modules = torch.nn.ModuleList([
|
||||
transformer.transformer_blocks[i].attn1.processor for i in range(len(transformer.transformer_blocks))
|
||||
] + [
|
||||
transformer.transformer_blocks[i].attn2.processor for i in range(len(transformer.transformer_blocks))
|
||||
])
|
||||
else:
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
self.adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
|
||||
# make a getter to see if is active
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def forward(self, input):
|
||||
return input
|
||||
256
toolkit/models/size_agnostic_feature_encoder.py
Normal file
256
toolkit/models/size_agnostic_feature_encoder.py
Normal file
@@ -0,0 +1,256 @@
|
||||
import os
|
||||
from typing import Union, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers.image_processing_utils import BaseImageProcessor
|
||||
|
||||
|
||||
class SAFEReducerBlock(nn.Module):
|
||||
"""
|
||||
This is the block that reduces the size of an vactor w and h be half. It is designed to be iterative
|
||||
So it is run multiple times to reduce an image to a desired dimension while carrying a shrinking residual
|
||||
along for the ride. This is done to preserve information.
|
||||
"""
|
||||
def __init__(self, channels=512):
|
||||
super(SAFEReducerBlock, self).__init__()
|
||||
self.channels = channels
|
||||
|
||||
activation = nn.GELU
|
||||
|
||||
self.reducer = nn.Sequential(
|
||||
nn.Conv2d(channels, channels, kernel_size=3, padding=1),
|
||||
activation(),
|
||||
nn.BatchNorm2d(channels),
|
||||
nn.Conv2d(channels, channels, kernel_size=3, padding=1),
|
||||
activation(),
|
||||
nn.BatchNorm2d(channels),
|
||||
nn.AvgPool2d(kernel_size=2, stride=2),
|
||||
)
|
||||
self.residual_shrink = nn.AvgPool2d(kernel_size=2, stride=2)
|
||||
|
||||
def forward(self, x):
|
||||
res = self.residual_shrink(x)
|
||||
reduced = self.reducer(x)
|
||||
return reduced + res
|
||||
|
||||
|
||||
class SizeAgnosticFeatureEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=3,
|
||||
num_tokens=8,
|
||||
num_vectors=768,
|
||||
reducer_channels=512,
|
||||
channels=2048,
|
||||
downscale_factor: int = 8,
|
||||
):
|
||||
super(SizeAgnosticFeatureEncoder, self).__init__()
|
||||
self.num_tokens = num_tokens
|
||||
self.num_vectors = num_vectors
|
||||
self.channels = channels
|
||||
self.reducer_channels = reducer_channels
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# input is minimum of (bs, 3, 256, 256)
|
||||
|
||||
subpixel_channels = in_channels * downscale_factor ** 2
|
||||
|
||||
# PixelUnshuffle(8 = # (bs, 3, 32, 32) -> (bs, 192, 32, 32)
|
||||
# PixelUnshuffle(16 = # (bs, 3, 16, 16) -> (bs, 48, 16, 16)
|
||||
|
||||
self.unshuffle = nn.PixelUnshuffle(downscale_factor) # (bs, 3, 256, 256) -> (bs, 192, 32, 32)
|
||||
|
||||
self.conv_in = nn.Conv2d(subpixel_channels, reducer_channels, kernel_size=3, padding=1) # (bs, 192, 32, 32) -> (bs, 512, 32, 32)
|
||||
|
||||
# run as many times as needed to get to min feature of 8 on the smallest dimension
|
||||
self.reducer = SAFEReducerBlock(reducer_channels) # (bs, 512, 32, 32) -> (bs, 512, 8, 8)
|
||||
|
||||
self.reduced_out = nn.Conv2d(
|
||||
reducer_channels, self.channels, kernel_size=3, padding=1
|
||||
) # (bs, 512, 8, 8) -> (bs, 2048, 8, 8)
|
||||
|
||||
# (bs, 2048, 8, 8)
|
||||
self.block1 = SAFEReducerBlock(self.channels) # (bs, 2048, 8, 8) -> (bs, 2048, 4, 4)
|
||||
self.block2 = SAFEReducerBlock(self.channels) # (bs, 2048, 8, 8) -> (bs, 2048, 2, 2)
|
||||
|
||||
# reduce mean of dims 2 and 3
|
||||
self.adaptive_pool = nn.Sequential(
|
||||
nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Flatten(),
|
||||
)
|
||||
|
||||
# (bs, 2048)
|
||||
# linear layer to (bs, self.num_vectors * self.num_tokens)
|
||||
self.fc1 = nn.Linear(self.channels, self.num_vectors * self.num_tokens)
|
||||
|
||||
# (bs, self.num_vectors * self.num_tokens) = (bs, 8 * 768) = (bs, 6144)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.unshuffle(x)
|
||||
x = self.conv_in(x)
|
||||
|
||||
while True:
|
||||
# reduce until we get as close to 8x8 as possible without going under
|
||||
x = self.reducer(x)
|
||||
if x.shape[2] // 2 < 8 or x.shape[3] // 2 < 8:
|
||||
break
|
||||
|
||||
x = self.reduced_out(x)
|
||||
x = self.block1(x)
|
||||
x = self.block2(x)
|
||||
x = self.adaptive_pool(x)
|
||||
x = self.fc1(x)
|
||||
|
||||
# reshape
|
||||
x = x.view(-1, self.num_tokens, self.num_vectors)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SAFEIPReturn:
|
||||
def __init__(self, pixel_values):
|
||||
self.pixel_values = pixel_values
|
||||
|
||||
|
||||
class SAFEImageProcessor(BaseImageProcessor):
|
||||
def __init__(
|
||||
self,
|
||||
max_size=1024,
|
||||
min_size=256,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.max_size = max_size
|
||||
self.min_size = min_size
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
pretrained_model_name_or_path: Union[str, os.PathLike],
|
||||
cache_dir: Optional[Union[str, os.PathLike]] = None,
|
||||
force_download: bool = False,
|
||||
local_files_only: bool = False,
|
||||
token: Optional[Union[str, bool]] = None,
|
||||
revision: str = "main",
|
||||
**kwargs,
|
||||
):
|
||||
# not needed
|
||||
return cls(**kwargs)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
images,
|
||||
**kwargs
|
||||
):
|
||||
# TODO allow for random resizing
|
||||
# comes in 0 - 1 range
|
||||
# if any size is smaller than 256, resize to 256
|
||||
# if any size is larger than max_size, resize to max_size
|
||||
if images.min() < -0.3 or images.max() > 1.3:
|
||||
raise ValueError(
|
||||
"images fed into SAFEImageProcessor values must be between 0 and 1. Got min: {}, max: {}".format(
|
||||
images.min(), images.max()
|
||||
))
|
||||
|
||||
# make sure we have (bs, 3, h, w)
|
||||
while len(images.shape) < 4:
|
||||
images = images.unsqueeze(0)
|
||||
|
||||
# expand to 3 channels if we only have 1 channel
|
||||
if images.shape[1] == 1:
|
||||
images = torch.cat([images, images, images], dim=1)
|
||||
|
||||
width = images.shape[3]
|
||||
height = images.shape[2]
|
||||
|
||||
if width < self.min_size or height < self.min_size:
|
||||
# scale up so that the smallest size is 256
|
||||
if width < height:
|
||||
new_width = self.min_size
|
||||
new_height = int(height * (self.min_size / width))
|
||||
else:
|
||||
new_height = self.min_size
|
||||
new_width = int(width * (self.min_size / height))
|
||||
images = nn.functional.interpolate(images, size=(new_height, new_width), mode='bilinear',
|
||||
align_corners=False)
|
||||
|
||||
elif width > self.max_size or height > self.max_size:
|
||||
# scale down so that the largest size is max_size but do not shrink the other size below 256
|
||||
if width > height:
|
||||
new_width = self.max_size
|
||||
new_height = int(height * (self.max_size / width))
|
||||
else:
|
||||
new_height = self.max_size
|
||||
new_width = int(width * (self.max_size / height))
|
||||
|
||||
if new_width < self.min_size:
|
||||
new_width = self.min_size
|
||||
new_height = int(height * (self.min_size / width))
|
||||
|
||||
if new_height < self.min_size:
|
||||
new_height = self.min_size
|
||||
new_width = int(width * (self.min_size / height))
|
||||
|
||||
images = nn.functional.interpolate(images, size=(new_height, new_width), mode='bilinear',
|
||||
align_corners=False)
|
||||
|
||||
# if wither side is not divisible by 16, mirror pad to make it so
|
||||
if images.shape[2] % 16 != 0:
|
||||
pad = 16 - (images.shape[2] % 16)
|
||||
pad1 = pad // 2
|
||||
pad2 = pad - pad1
|
||||
images = nn.functional.pad(images, (0, 0, pad1, pad2), mode='reflect')
|
||||
if images.shape[3] % 16 != 0:
|
||||
pad = 16 - (images.shape[3] % 16)
|
||||
pad1 = pad // 2
|
||||
pad2 = pad - pad1
|
||||
images = nn.functional.pad(images, (pad1, pad2, 0, 0), mode='reflect')
|
||||
|
||||
return SAFEIPReturn(images)
|
||||
|
||||
|
||||
class SAFEVMConfig:
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=3,
|
||||
num_tokens=8,
|
||||
num_vectors=768,
|
||||
reducer_channels=512,
|
||||
channels=2048,
|
||||
downscale_factor: int = 8,
|
||||
**kwargs
|
||||
):
|
||||
self.in_channels = in_channels
|
||||
self.num_tokens = num_tokens
|
||||
self.num_vectors = num_vectors
|
||||
self.reducer_channels = reducer_channels
|
||||
self.channels = channels
|
||||
self.downscale_factor = downscale_factor
|
||||
self.image_size = 224
|
||||
|
||||
self.hidden_size = num_vectors
|
||||
self.projection_dim = num_vectors
|
||||
|
||||
|
||||
class SAFEVMReturn:
|
||||
def __init__(self, output):
|
||||
self.output = output
|
||||
# todo actually do hidden states. This is just for code compatability for now
|
||||
self.hidden_states = [output for _ in range(13)]
|
||||
|
||||
|
||||
class SAFEVisionModel(SizeAgnosticFeatureEncoder):
|
||||
def __init__(self, **kwargs):
|
||||
self.config = SAFEVMConfig(**kwargs)
|
||||
self.image_size = None
|
||||
# super().__init__(**kwargs)
|
||||
super(SAFEVisionModel, self).__init__(**kwargs)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, *args, **kwargs):
|
||||
# not needed
|
||||
return SAFEVisionModel(**kwargs)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
return SAFEVMReturn(super().forward(x))
|
||||
460
toolkit/models/te_adapter.py
Normal file
460
toolkit/models/te_adapter.py
Normal file
@@ -0,0 +1,460 @@
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Union, TYPE_CHECKING
|
||||
|
||||
|
||||
from transformers import T5EncoderModel, CLIPTextModel, CLIPTokenizer, T5Tokenizer, CLIPTextModelWithProjection
|
||||
from diffusers.models.embeddings import PixArtAlphaTextProjection
|
||||
|
||||
from toolkit import train_tools
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from diffusers import Transformer2DModel
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
from ipadapter.ip_adapter.attention_processor import AttnProcessor2_0
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PixArtSigmaPipeline
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
class TEAdapterCaptionProjection(nn.Module):
|
||||
def __init__(self, caption_channels, adapter: 'TEAdapter'):
|
||||
super().__init__()
|
||||
in_features = caption_channels
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
sd = adapter.sd_ref()
|
||||
self.parent_module_ref = weakref.ref(sd.unet.caption_projection)
|
||||
parent_module = self.parent_module_ref()
|
||||
self.linear_1 = nn.Linear(
|
||||
in_features=in_features,
|
||||
out_features=parent_module.linear_1.out_features,
|
||||
bias=True
|
||||
)
|
||||
self.linear_2 = nn.Linear(
|
||||
in_features=parent_module.linear_2.in_features,
|
||||
out_features=parent_module.linear_2.out_features,
|
||||
bias=True
|
||||
)
|
||||
|
||||
# save the orig forward
|
||||
parent_module.linear_1.orig_forward = parent_module.linear_1.forward
|
||||
parent_module.linear_2.orig_forward = parent_module.linear_2.forward
|
||||
|
||||
# replace original forward
|
||||
parent_module.orig_forward = parent_module.forward
|
||||
parent_module.forward = self.forward
|
||||
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
@property
|
||||
def unconditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().unconditional_embeds
|
||||
|
||||
@property
|
||||
def conditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().conditional_embeds
|
||||
|
||||
def forward(self, caption):
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
adapter_hidden_states = self.conditional_embeds.text_embeds
|
||||
# check if we are doing unconditional
|
||||
if self.unconditional_embeds is not None and adapter_hidden_states.shape[0] != caption.shape[0]:
|
||||
# concat unconditional to match the hidden state batch size
|
||||
if self.unconditional_embeds.text_embeds.shape[0] == 1 and adapter_hidden_states.shape[0] != 1:
|
||||
unconditional = torch.cat([self.unconditional_embeds.text_embeds] * adapter_hidden_states.shape[0], dim=0)
|
||||
else:
|
||||
unconditional = self.unconditional_embeds.text_embeds
|
||||
adapter_hidden_states = torch.cat([unconditional, adapter_hidden_states], dim=0)
|
||||
hidden_states = self.linear_1(adapter_hidden_states)
|
||||
hidden_states = self.parent_module_ref().act_1(hidden_states)
|
||||
hidden_states = self.linear_2(hidden_states)
|
||||
return hidden_states
|
||||
else:
|
||||
return self.parent_module_ref().orig_forward(caption)
|
||||
|
||||
|
||||
class TEAdapterAttnProcessor(nn.Module):
|
||||
r"""
|
||||
Attention processor for Custom TE for PyTorch 2.0.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
scale (`float`, defaults to 1.0):
|
||||
the weight scale of image prompt.
|
||||
num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16):
|
||||
The context length of the image features.
|
||||
adapter
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4, adapter=None,
|
||||
adapter_hidden_size=None, layer_name=None):
|
||||
super().__init__()
|
||||
self.layer_name = layer_name
|
||||
|
||||
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.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.adapter_hidden_size = adapter_hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
self.to_k_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=False)
|
||||
self.to_v_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=False)
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
@property
|
||||
def unconditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().unconditional_embeds
|
||||
|
||||
@property
|
||||
def conditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().conditional_embeds
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
is_active = self.adapter_ref().is_active
|
||||
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)
|
||||
|
||||
# will be none if disabled
|
||||
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)
|
||||
|
||||
# only use one TE or the other. If our adapter is active only use ours
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
adapter_hidden_states = self.conditional_embeds.text_embeds
|
||||
# check if we are doing unconditional
|
||||
if self.unconditional_embeds is not None and adapter_hidden_states.shape[0] != encoder_hidden_states.shape[0]:
|
||||
# concat unconditional to match the hidden state batch size
|
||||
if self.unconditional_embeds.text_embeds.shape[0] == 1 and adapter_hidden_states.shape[0] != 1:
|
||||
unconditional = torch.cat([self.unconditional_embeds.text_embeds] * adapter_hidden_states.shape[0], dim=0)
|
||||
else:
|
||||
unconditional = self.unconditional_embeds.text_embeds
|
||||
adapter_hidden_states = torch.cat([unconditional, adapter_hidden_states], dim=0)
|
||||
# for ip-adapter
|
||||
key = self.to_k_adapter(adapter_hidden_states)
|
||||
value = self.to_v_adapter(adapter_hidden_states)
|
||||
else:
|
||||
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)
|
||||
|
||||
try:
|
||||
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)
|
||||
except RuntimeError:
|
||||
raise RuntimeError(f"key shape: {key.shape}, value shape: {value.shape}")
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
# remove attn mask if doing clip
|
||||
if self.adapter_ref().adapter_ref().config.text_encoder_arch == "clip":
|
||||
attention_mask = None
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
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 TEAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
te: Union[T5EncoderModel],
|
||||
tokenizer: CLIPTokenizer
|
||||
):
|
||||
super(TEAdapter, self).__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.te_ref: weakref.ref = weakref.ref(te)
|
||||
self.tokenizer_ref: weakref.ref = weakref.ref(tokenizer)
|
||||
self.adapter_modules = []
|
||||
self.caption_projection = None
|
||||
self.embeds_store = []
|
||||
is_pixart = sd.is_pixart
|
||||
|
||||
if self.adapter_ref().config.text_encoder_arch == "t5" or self.adapter_ref().config.text_encoder_arch == "pile-t5":
|
||||
self.token_size = self.te_ref().config.d_model
|
||||
else:
|
||||
self.token_size = self.te_ref().config.hidden_size
|
||||
|
||||
# add text projection if is sdxl
|
||||
self.text_projection = None
|
||||
if sd.is_xl:
|
||||
clip_with_projection: CLIPTextModelWithProjection = sd.text_encoder[0]
|
||||
self.text_projection = nn.Linear(te.config.hidden_size, clip_with_projection.config.projection_dim, bias=False)
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
attn_dict_map = {
|
||||
|
||||
}
|
||||
module_idx = 0
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
attn_processor_keys = []
|
||||
if is_pixart:
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
|
||||
|
||||
# cross attention
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
|
||||
|
||||
else:
|
||||
attn_processor_keys = list(sd.unet.attn_processors.keys())
|
||||
|
||||
attn_processor_names = []
|
||||
|
||||
blocks = []
|
||||
transformer_blocks = []
|
||||
for name in attn_processor_keys:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") or name.endswith("attn1") else \
|
||||
sd.unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = sd.unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(sd.unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
hidden_size = sd.unet.config['cross_attention_dim']
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AttnProcessor2_0()
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
|
||||
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
|
||||
|
||||
# add zero padding to the adapter
|
||||
if to_k_adapter.shape[1] < self.token_size:
|
||||
to_k_adapter = torch.cat([
|
||||
to_k_adapter,
|
||||
torch.randn(to_k_adapter.shape[0], self.token_size - to_k_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
to_v_adapter = torch.cat([
|
||||
to_v_adapter,
|
||||
torch.randn(to_v_adapter.shape[0], self.token_size - to_v_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
elif to_k_adapter.shape[1] > self.token_size:
|
||||
to_k_adapter = to_k_adapter[:, :self.token_size]
|
||||
to_v_adapter = to_v_adapter[:, :self.token_size]
|
||||
else:
|
||||
to_k_adapter = to_k_adapter
|
||||
to_v_adapter = to_v_adapter
|
||||
|
||||
# todo resize to the TE hidden size
|
||||
weights = {
|
||||
"to_k_adapter.weight": to_k_adapter,
|
||||
"to_v_adapter.weight": to_v_adapter,
|
||||
}
|
||||
|
||||
if self.sd_ref().is_pixart:
|
||||
# pixart is much more sensitive
|
||||
weights = {
|
||||
"to_k_adapter.weight": weights["to_k_adapter.weight"] * 0.01,
|
||||
"to_v_adapter.weight": weights["to_v_adapter.weight"] * 0.01,
|
||||
}
|
||||
|
||||
attn_procs[name] = TEAdapterAttnProcessor(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
num_tokens=self.adapter_ref().config.num_tokens,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
layer_name=layer_name
|
||||
)
|
||||
attn_procs[name].load_state_dict(weights)
|
||||
self.adapter_modules.append(attn_procs[name])
|
||||
if self.sd_ref().is_pixart:
|
||||
# we have to set them ourselves
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn1.processor = attn_procs[f"transformer_blocks.{i}.attn1"]
|
||||
module.attn2.processor = attn_procs[f"transformer_blocks.{i}.attn2"]
|
||||
self.adapter_modules = torch.nn.ModuleList(
|
||||
[
|
||||
transformer.transformer_blocks[i].attn2.processor for i in
|
||||
range(len(transformer.transformer_blocks))
|
||||
])
|
||||
self.caption_projection = TEAdapterCaptionProjection(
|
||||
caption_channels=self.token_size,
|
||||
adapter=self,
|
||||
)
|
||||
|
||||
else:
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
self.adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
|
||||
# make a getter to see if is active
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def encode_text(self, text):
|
||||
te: T5EncoderModel = self.te_ref()
|
||||
tokenizer: T5Tokenizer = self.tokenizer_ref()
|
||||
attn_mask_float = None
|
||||
|
||||
# input_ids = tokenizer(
|
||||
# text,
|
||||
# max_length=77,
|
||||
# padding="max_length",
|
||||
# truncation=True,
|
||||
# return_tensors="pt",
|
||||
# ).input_ids.to(te.device)
|
||||
# outputs = te(input_ids=input_ids)
|
||||
# outputs = outputs.last_hidden_state
|
||||
if self.adapter_ref().config.text_encoder_arch == "clip":
|
||||
embeds = train_tools.encode_prompts(
|
||||
tokenizer,
|
||||
te,
|
||||
text,
|
||||
truncate=True,
|
||||
max_length=self.adapter_ref().config.num_tokens,
|
||||
)
|
||||
attention_mask = torch.ones(embeds.shape[:2], device=embeds.device)
|
||||
|
||||
elif self.adapter_ref().config.text_encoder_arch == "pile-t5":
|
||||
# just use aura pile
|
||||
embeds, attention_mask = train_tools.encode_prompts_auraflow(
|
||||
tokenizer,
|
||||
te,
|
||||
text,
|
||||
truncate=True,
|
||||
max_length=self.adapter_ref().config.num_tokens,
|
||||
)
|
||||
|
||||
else:
|
||||
embeds, attention_mask = train_tools.encode_prompts_pixart(
|
||||
tokenizer,
|
||||
te,
|
||||
text,
|
||||
truncate=True,
|
||||
max_length=self.adapter_ref().config.num_tokens,
|
||||
)
|
||||
if attention_mask is not None:
|
||||
attn_mask_float = attention_mask.to(embeds.device, dtype=embeds.dtype)
|
||||
if self.text_projection is not None:
|
||||
# pool the output of embeds ignoring 0 in the attention mask
|
||||
if attn_mask_float is not None:
|
||||
pooled_output = embeds * attn_mask_float.unsqueeze(-1)
|
||||
else:
|
||||
pooled_output = embeds
|
||||
|
||||
# reduce along dim 1 while maintaining batch and dim 2
|
||||
pooled_output_sum = pooled_output.sum(dim=1)
|
||||
|
||||
if attn_mask_float is not None:
|
||||
attn_mask_sum = attn_mask_float.sum(dim=1).unsqueeze(-1)
|
||||
|
||||
pooled_output = pooled_output_sum / attn_mask_sum
|
||||
|
||||
pooled_embeds = self.text_projection(pooled_output)
|
||||
|
||||
prompt_embeds = PromptEmbeds(
|
||||
(embeds, pooled_embeds),
|
||||
attention_mask=attention_mask,
|
||||
).detach()
|
||||
|
||||
else:
|
||||
|
||||
prompt_embeds = PromptEmbeds(
|
||||
embeds,
|
||||
attention_mask=attention_mask,
|
||||
).detach()
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
|
||||
|
||||
def forward(self, input):
|
||||
return input
|
||||
253
toolkit/models/te_aug_adapter.py
Normal file
253
toolkit/models/te_aug_adapter.py
Normal file
@@ -0,0 +1,253 @@
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Union, TYPE_CHECKING, Optional, Tuple
|
||||
|
||||
from transformers import T5EncoderModel, CLIPTextModel, CLIPTokenizer, T5Tokenizer
|
||||
from transformers.models.clip.modeling_clip import CLIPEncoder, CLIPAttention
|
||||
|
||||
from toolkit.models.zipper_resampler import ZipperResampler, ZipperModule
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.resampler import Resampler
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
from ipadapter.ip_adapter.attention_processor import AttnProcessor2_0
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
class TEAugAdapterCLIPAttention(nn.Module):
|
||||
"""Multi-headed attention from 'Attention Is All You Need' paper"""
|
||||
|
||||
def __init__(self, attn_module: 'CLIPAttention', adapter: 'TEAugAdapter'):
|
||||
super().__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.attn_module_ref: weakref.ref = weakref.ref(attn_module)
|
||||
self.k_proj_adapter = nn.Linear(attn_module.embed_dim, attn_module.embed_dim)
|
||||
self.v_proj_adapter = nn.Linear(attn_module.embed_dim, attn_module.embed_dim)
|
||||
# copy the weights from the original module
|
||||
self.k_proj_adapter.weight.data = attn_module.k_proj.weight.data.clone() * 0.01
|
||||
self.v_proj_adapter.weight.data = attn_module.v_proj.weight.data.clone() * 0.01
|
||||
#reset the bias
|
||||
self.k_proj_adapter.bias.data = attn_module.k_proj.bias.data.clone() * 0.001
|
||||
self.v_proj_adapter.bias.data = attn_module.v_proj.bias.data.clone() * 0.001
|
||||
|
||||
self.zipper = ZipperModule(
|
||||
in_size=attn_module.embed_dim,
|
||||
in_tokens=77 * 2,
|
||||
out_size=attn_module.embed_dim,
|
||||
out_tokens=77,
|
||||
hidden_size=attn_module.embed_dim,
|
||||
hidden_tokens=77,
|
||||
)
|
||||
# self.k_proj_adapter.weight.data = torch.zeros_like(attn_module.k_proj.weight.data)
|
||||
# self.v_proj_adapter.weight.data = torch.zeros_like(attn_module.v_proj.weight.data)
|
||||
# #reset the bias
|
||||
# self.k_proj_adapter.bias.data = torch.zeros_like(attn_module.k_proj.bias.data)
|
||||
# self.v_proj_adapter.bias.data = torch.zeros_like(attn_module.v_proj.bias.data)
|
||||
|
||||
# replace the original forward with our forward
|
||||
self.original_forward = attn_module.forward
|
||||
attn_module.forward = self.forward
|
||||
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
||||
return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
causal_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = False,
|
||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
|
||||
"""Input shape: Batch x Time x Channel"""
|
||||
|
||||
attn_module = self.attn_module_ref()
|
||||
|
||||
bsz, tgt_len, embed_dim = hidden_states.size()
|
||||
|
||||
# get query proj
|
||||
query_states = attn_module.q_proj(hidden_states) * attn_module.scale
|
||||
key_states = attn_module._shape(attn_module.k_proj(hidden_states), -1, bsz)
|
||||
value_states = attn_module._shape(attn_module.v_proj(hidden_states), -1, bsz)
|
||||
|
||||
proj_shape = (bsz * attn_module.num_heads, -1, attn_module.head_dim)
|
||||
query_states = attn_module._shape(query_states, tgt_len, bsz).view(*proj_shape)
|
||||
key_states = key_states.view(*proj_shape)
|
||||
value_states = value_states.view(*proj_shape)
|
||||
|
||||
src_len = key_states.size(1)
|
||||
attn_weights = torch.bmm(query_states, key_states.transpose(1, 2))
|
||||
|
||||
if attn_weights.size() != (bsz * attn_module.num_heads, tgt_len, src_len):
|
||||
raise ValueError(
|
||||
f"Attention weights should be of size {(bsz * attn_module.num_heads, tgt_len, src_len)}, but is"
|
||||
f" {attn_weights.size()}"
|
||||
)
|
||||
|
||||
# apply the causal_attention_mask first
|
||||
if causal_attention_mask is not None:
|
||||
if causal_attention_mask.size() != (bsz, 1, tgt_len, src_len):
|
||||
raise ValueError(
|
||||
f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is"
|
||||
f" {causal_attention_mask.size()}"
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz, attn_module.num_heads, tgt_len, src_len) + causal_attention_mask
|
||||
attn_weights = attn_weights.view(bsz * attn_module.num_heads, tgt_len, src_len)
|
||||
|
||||
if attention_mask is not None:
|
||||
if attention_mask.size() != (bsz, 1, tgt_len, src_len):
|
||||
raise ValueError(
|
||||
f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is {attention_mask.size()}"
|
||||
)
|
||||
attn_weights = attn_weights.view(bsz, attn_module.num_heads, tgt_len, src_len) + attention_mask
|
||||
attn_weights = attn_weights.view(bsz * attn_module.num_heads, tgt_len, src_len)
|
||||
|
||||
attn_weights = nn.functional.softmax(attn_weights, dim=-1)
|
||||
|
||||
if output_attentions:
|
||||
# this operation is a bit akward, but it's required to
|
||||
# make sure that attn_weights keeps its gradient.
|
||||
# In order to do so, attn_weights have to reshaped
|
||||
# twice and have to be reused in the following
|
||||
attn_weights_reshaped = attn_weights.view(bsz, attn_module.num_heads, tgt_len, src_len)
|
||||
attn_weights = attn_weights_reshaped.view(bsz * attn_module.num_heads, tgt_len, src_len)
|
||||
else:
|
||||
attn_weights_reshaped = None
|
||||
|
||||
attn_probs = nn.functional.dropout(attn_weights, p=attn_module.dropout, training=self.training)
|
||||
|
||||
attn_output = torch.bmm(attn_probs, value_states)
|
||||
|
||||
if attn_output.size() != (bsz * attn_module.num_heads, tgt_len, attn_module.head_dim):
|
||||
raise ValueError(
|
||||
f"`attn_output` should be of size {(bsz, attn_module.num_heads, tgt_len, attn_module.head_dim)}, but is"
|
||||
f" {attn_output.size()}"
|
||||
)
|
||||
|
||||
attn_output = attn_output.view(bsz, attn_module.num_heads, tgt_len, attn_module.head_dim)
|
||||
attn_output = attn_output.transpose(1, 2)
|
||||
attn_output = attn_output.reshape(bsz, tgt_len, embed_dim)
|
||||
|
||||
adapter: 'CustomAdapter' = self.adapter_ref().adapter_ref()
|
||||
if self.adapter_ref().is_active and adapter.conditional_embeds is not None:
|
||||
# apply the adapter
|
||||
|
||||
if adapter.is_unconditional_run:
|
||||
embeds = adapter.unconditional_embeds
|
||||
else:
|
||||
embeds = adapter.conditional_embeds
|
||||
# if the shape is not the same on batch, we are doing cfg and need to concat unconditional as well
|
||||
if embeds.size(0) != bsz:
|
||||
embeds = torch.cat([adapter.unconditional_embeds, embeds], dim=0)
|
||||
|
||||
key_states_raw = self.k_proj_adapter(embeds)
|
||||
key_states = attn_module._shape(key_states_raw, -1, bsz)
|
||||
value_states_raw = self.v_proj_adapter(embeds)
|
||||
value_states = attn_module._shape(value_states_raw, -1, bsz)
|
||||
key_states = key_states.view(*proj_shape)
|
||||
value_states = value_states.view(*proj_shape)
|
||||
attn_weights = torch.bmm(query_states, key_states.transpose(1, 2))
|
||||
|
||||
attn_weights = nn.functional.softmax(attn_weights, dim=-1)
|
||||
attn_probs = nn.functional.dropout(attn_weights, p=attn_module.dropout, training=self.training)
|
||||
attn_output_adapter = torch.bmm(attn_probs, value_states)
|
||||
|
||||
if attn_output_adapter.size() != (bsz * attn_module.num_heads, tgt_len, attn_module.head_dim):
|
||||
raise ValueError(
|
||||
f"`attn_output_adapter` should be of size {(bsz, attn_module.num_heads, tgt_len, attn_module.head_dim)}, but is"
|
||||
f" {attn_output_adapter.size()}"
|
||||
)
|
||||
|
||||
attn_output_adapter = attn_output_adapter.view(bsz, attn_module.num_heads, tgt_len, attn_module.head_dim)
|
||||
attn_output_adapter = attn_output_adapter.transpose(1, 2)
|
||||
attn_output_adapter = attn_output_adapter.reshape(bsz, tgt_len, embed_dim)
|
||||
|
||||
attn_output_adapter = self.zipper(torch.cat([attn_output_adapter, attn_output], dim=1))
|
||||
|
||||
# attn_output_adapter = attn_module.out_proj(attn_output_adapter)
|
||||
attn_output = attn_output + attn_output_adapter
|
||||
|
||||
attn_output = attn_module.out_proj(attn_output)
|
||||
|
||||
return attn_output, attn_weights_reshaped
|
||||
|
||||
class TEAugAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
):
|
||||
super(TEAugAdapter, self).__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
|
||||
if isinstance(sd.text_encoder, list):
|
||||
raise ValueError("Dual text encoders is not yet supported")
|
||||
|
||||
# dim will come from text encoder
|
||||
# dim = sd.unet.config['cross_attention_dim']
|
||||
text_encoder: CLIPTextModel = sd.text_encoder
|
||||
dim = text_encoder.config.hidden_size
|
||||
|
||||
clip_encoder: CLIPEncoder = text_encoder.text_model.encoder
|
||||
# dim = clip_encoder.layers[-1].self_attn
|
||||
|
||||
if hasattr(adapter.vision_encoder.config, 'hidden_sizes'):
|
||||
embedding_dim = adapter.vision_encoder.config.hidden_sizes[-1]
|
||||
else:
|
||||
embedding_dim = adapter.vision_encoder.config.hidden_size
|
||||
|
||||
image_encoder_state_dict = adapter.vision_encoder.state_dict()
|
||||
# max_seq_len = CLIP tokens + CLS token
|
||||
in_tokens = 257
|
||||
if "vision_model.embeddings.position_embedding.weight" in image_encoder_state_dict:
|
||||
# clip
|
||||
in_tokens = int(image_encoder_state_dict["vision_model.embeddings.position_embedding.weight"].shape[0])
|
||||
|
||||
if adapter.config.image_encoder_arch.startswith('convnext'):
|
||||
in_tokens = 16 * 16
|
||||
embedding_dim = adapter.vision_encoder.config.hidden_sizes[-1]
|
||||
|
||||
out_tokens = adapter.config.num_tokens if adapter.config.num_tokens > 0 else in_tokens
|
||||
self.image_proj_model = ZipperModule(
|
||||
in_size=embedding_dim,
|
||||
in_tokens=in_tokens,
|
||||
out_size=dim,
|
||||
out_tokens=out_tokens,
|
||||
hidden_size=dim,
|
||||
hidden_tokens=out_tokens,
|
||||
)
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
for idx, layer in enumerate(clip_encoder.layers):
|
||||
name = f"clip_attention.{idx}"
|
||||
attn_procs[name] = TEAugAdapterCLIPAttention(
|
||||
layer.self_attn,
|
||||
self
|
||||
)
|
||||
|
||||
self.adapter_modules = torch.nn.ModuleList(list(attn_procs.values()))
|
||||
|
||||
# make a getter to see if is active
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
|
||||
def forward(self, input):
|
||||
# # apply the adapter
|
||||
input = self.image_proj_model(input)
|
||||
# self.embeds = input
|
||||
return input
|
||||
1437
toolkit/models/unified_training_model.py
Normal file
1437
toolkit/models/unified_training_model.py
Normal file
File diff suppressed because it is too large
Load Diff
623
toolkit/models/vd_adapter.py
Normal file
623
toolkit/models/vd_adapter.py
Normal file
@@ -0,0 +1,623 @@
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Union, TYPE_CHECKING, Optional
|
||||
|
||||
from diffusers import Transformer2DModel, FluxTransformer2DModel
|
||||
from transformers import T5EncoderModel, CLIPTextModel, CLIPTokenizer, T5Tokenizer, CLIPVisionModelWithProjection
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
class AttnProcessor2_0(torch.nn.Module):
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size=None,
|
||||
cross_attention_dim=None,
|
||||
):
|
||||
super().__init__()
|
||||
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.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
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)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
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 VisionDirectAdapterAttnProcessor(nn.Module):
|
||||
r"""
|
||||
Attention processor for Custom TE for PyTorch 2.0.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
scale (`float`, defaults to 1.0):
|
||||
the weight scale of image prompt.
|
||||
adapter
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, adapter=None,
|
||||
adapter_hidden_size=None, has_bias=False, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
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.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.adapter_hidden_size = adapter_hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
|
||||
self.to_k_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
self.to_v_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
# return False
|
||||
|
||||
@property
|
||||
def unconditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().unconditional_embeds
|
||||
|
||||
@property
|
||||
def conditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().conditional_embeds
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
is_active = self.adapter_ref().is_active
|
||||
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)
|
||||
|
||||
# will be none if disabled
|
||||
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
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# only use one TE or the other. If our adapter is active only use ours
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
|
||||
adapter_hidden_states = self.conditional_embeds
|
||||
if adapter_hidden_states.shape[0] < batch_size:
|
||||
adapter_hidden_states = torch.cat([
|
||||
self.unconditional_embeds,
|
||||
adapter_hidden_states
|
||||
], dim=0)
|
||||
# if it is image embeds, we need to add a 1 dim at inx 1
|
||||
if len(adapter_hidden_states.shape) == 2:
|
||||
adapter_hidden_states = adapter_hidden_states.unsqueeze(1)
|
||||
# conditional_batch_size = adapter_hidden_states.shape[0]
|
||||
# conditional_query = query
|
||||
|
||||
# for ip-adapter
|
||||
vd_key = self.to_k_adapter(adapter_hidden_states)
|
||||
vd_value = self.to_v_adapter(adapter_hidden_states)
|
||||
|
||||
vd_key = vd_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
vd_value = vd_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
|
||||
vd_hidden_states = F.scaled_dot_product_attention(
|
||||
query, vd_key, vd_value, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
vd_hidden_states = vd_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
vd_hidden_states = vd_hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states = hidden_states + self.scale * vd_hidden_states
|
||||
|
||||
|
||||
# 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 CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
"""Attention processor used typically in processing the SD3-like self-attention projections."""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, adapter=None,
|
||||
adapter_hidden_size=None, has_bias=False, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
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.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.adapter_hidden_size = adapter_hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
|
||||
self.to_k_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
self.to_v_adapter = nn.Linear(adapter_hidden_size, hidden_size, bias=has_bias)
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
# return False
|
||||
|
||||
@property
|
||||
def unconditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().unconditional_embeds
|
||||
|
||||
@property
|
||||
def conditional_embeds(self):
|
||||
return self.adapter_ref().adapter_ref().conditional_embeds
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
is_active = self.adapter_ref().is_active
|
||||
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)
|
||||
context_input_ndim = encoder_hidden_states.ndim
|
||||
if context_input_ndim == 4:
|
||||
batch_size, channel, height, width = encoder_hidden_states.shape
|
||||
encoder_hidden_states = encoder_hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size = encoder_hidden_states.shape[0]
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(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)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# `context` projections.
|
||||
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
|
||||
|
||||
# attention
|
||||
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
|
||||
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
|
||||
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
# YiYi to-do: update uising apply_rotary_emb
|
||||
# from ..embeddings import apply_rotary_emb
|
||||
# query = apply_rotary_emb(query, image_rotary_emb)
|
||||
# key = apply_rotary_emb(key, image_rotary_emb)
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# do ip adapter
|
||||
# will be none if disabled
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
adapter_hidden_states = self.conditional_embeds
|
||||
if adapter_hidden_states.shape[0] < batch_size:
|
||||
adapter_hidden_states = torch.cat([
|
||||
self.unconditional_embeds,
|
||||
adapter_hidden_states
|
||||
], dim=0)
|
||||
# if it is image embeds, we need to add a 1 dim at inx 1
|
||||
if len(adapter_hidden_states.shape) == 2:
|
||||
adapter_hidden_states = adapter_hidden_states.unsqueeze(1)
|
||||
# conditional_batch_size = adapter_hidden_states.shape[0]
|
||||
# conditional_query = query
|
||||
|
||||
# for ip-adapter
|
||||
vd_key = self.to_k_adapter(adapter_hidden_states)
|
||||
vd_value = self.to_v_adapter(adapter_hidden_states)
|
||||
|
||||
vd_key = vd_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
vd_value = vd_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
|
||||
vd_hidden_states = F.scaled_dot_product_attention(
|
||||
query, vd_key, vd_value, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
vd_hidden_states = vd_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
vd_hidden_states = vd_hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states = hidden_states + self.scale * vd_hidden_states
|
||||
|
||||
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
if context_input_ndim == 4:
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class VisionDirectAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
vision_model: Union[CLIPVisionModelWithProjection],
|
||||
):
|
||||
super(VisionDirectAdapter, self).__init__()
|
||||
is_pixart = sd.is_pixart
|
||||
is_flux = sd.is_flux
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.vision_model_ref: weakref.ref = weakref.ref(vision_model)
|
||||
|
||||
if adapter.config.clip_layer == "image_embeds":
|
||||
self.token_size = vision_model.config.projection_dim
|
||||
else:
|
||||
self.token_size = vision_model.config.hidden_size
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
|
||||
attn_processor_keys = []
|
||||
if is_pixart:
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
|
||||
|
||||
# cross attention
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
|
||||
|
||||
elif is_flux:
|
||||
transformer: FluxTransformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn")
|
||||
|
||||
# single transformer blocks do not have cross attn
|
||||
# for i, module in transformer.single_transformer_blocks.named_children():
|
||||
# attn_processor_keys.append(f"single_transformer_blocks.{i}.attn")
|
||||
else:
|
||||
attn_processor_keys = list(sd.unet.attn_processors.keys())
|
||||
|
||||
for name in attn_processor_keys:
|
||||
if is_flux:
|
||||
cross_attention_dim = None
|
||||
else:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") else sd.unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = sd.unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(sd.unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
if is_flux:
|
||||
hidden_size = 3072
|
||||
else:
|
||||
hidden_size = sd.unet.config['cross_attention_dim']
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None and not is_flux:
|
||||
attn_procs[name] = AttnProcessor2_0()
|
||||
else:
|
||||
layer_name = name.split(".processor")[0]
|
||||
if f"{layer_name}.to_k.weight._data" in unet_sd and is_flux:
|
||||
# is quantized
|
||||
|
||||
to_k_adapter = torch.randn(hidden_size, hidden_size) * 0.01
|
||||
to_v_adapter = torch.randn(hidden_size, hidden_size) * 0.01
|
||||
to_k_adapter = to_k_adapter.to(self.sd_ref().torch_dtype)
|
||||
to_v_adapter = to_v_adapter.to(self.sd_ref().torch_dtype)
|
||||
else:
|
||||
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
|
||||
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
|
||||
|
||||
# add zero padding to the adapter
|
||||
if to_k_adapter.shape[1] < self.token_size:
|
||||
to_k_adapter = torch.cat([
|
||||
to_k_adapter,
|
||||
torch.randn(to_k_adapter.shape[0], self.token_size - to_k_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
to_v_adapter = torch.cat([
|
||||
to_v_adapter,
|
||||
torch.randn(to_v_adapter.shape[0], self.token_size - to_v_adapter.shape[1]).to(
|
||||
to_k_adapter.device, dtype=to_k_adapter.dtype) * 0.01
|
||||
],
|
||||
dim=1
|
||||
)
|
||||
elif to_k_adapter.shape[1] > self.token_size:
|
||||
to_k_adapter = to_k_adapter[:, :self.token_size]
|
||||
to_v_adapter = to_v_adapter[:, :self.token_size]
|
||||
# if is_pixart:
|
||||
# to_k_bias = to_k_bias[:self.token_size]
|
||||
# to_v_bias = to_v_bias[:self.token_size]
|
||||
else:
|
||||
to_k_adapter = to_k_adapter
|
||||
to_v_adapter = to_v_adapter
|
||||
# if is_pixart:
|
||||
# to_k_bias = to_k_bias
|
||||
# to_v_bias = to_v_bias
|
||||
|
||||
weights = {
|
||||
"to_k_adapter.weight": to_k_adapter * 0.01,
|
||||
"to_v_adapter.weight": to_v_adapter * 0.01,
|
||||
}
|
||||
# if is_pixart:
|
||||
# weights["to_k_adapter.bias"] = to_k_bias
|
||||
# weights["to_v_adapter.bias"] = to_v_bias\
|
||||
|
||||
if is_flux:
|
||||
attn_procs[name] = CustomFluxVDAttnProcessor2_0(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
has_bias=False,
|
||||
)
|
||||
else:
|
||||
attn_procs[name] = VisionDirectAdapterAttnProcessor(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
has_bias=False,
|
||||
)
|
||||
attn_procs[name].load_state_dict(weights)
|
||||
if self.sd_ref().is_pixart:
|
||||
# we have to set them ourselves
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn1.processor = attn_procs[f"transformer_blocks.{i}.attn1"]
|
||||
module.attn2.processor = attn_procs[f"transformer_blocks.{i}.attn2"]
|
||||
self.adapter_modules = torch.nn.ModuleList([
|
||||
transformer.transformer_blocks[i].attn1.processor for i in range(len(transformer.transformer_blocks))
|
||||
] + [
|
||||
transformer.transformer_blocks[i].attn2.processor for i in range(len(transformer.transformer_blocks))
|
||||
])
|
||||
elif self.sd_ref().is_flux:
|
||||
# we have to set them ourselves
|
||||
transformer: FluxTransformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"transformer_blocks.{i}.attn"]
|
||||
self.adapter_modules = torch.nn.ModuleList(
|
||||
[
|
||||
transformer.transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.transformer_blocks))
|
||||
])
|
||||
else:
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
self.adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
|
||||
# add the mlp layer
|
||||
self.mlp = MLP(
|
||||
in_dim=self.token_size,
|
||||
out_dim=self.token_size,
|
||||
hidden_dim=self.token_size,
|
||||
# dropout=0.1,
|
||||
use_residual=True
|
||||
)
|
||||
|
||||
# make a getter to see if is active
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def forward(self, input):
|
||||
return self.mlp(input)
|
||||
171
toolkit/models/zipper_resampler.py
Normal file
171
toolkit/models/zipper_resampler.py
Normal file
@@ -0,0 +1,171 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ContextualAlphaMask(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int = 768,
|
||||
):
|
||||
super(ContextualAlphaMask, self).__init__()
|
||||
self.dim = dim
|
||||
|
||||
half_dim = dim // 2
|
||||
quarter_dim = dim // 4
|
||||
|
||||
self.fc1 = nn.Linear(self.dim, self.dim)
|
||||
self.fc2 = nn.Linear(self.dim, half_dim)
|
||||
self.norm1 = nn.LayerNorm(half_dim)
|
||||
self.fc3 = nn.Linear(half_dim, half_dim)
|
||||
self.fc4 = nn.Linear(half_dim, quarter_dim)
|
||||
self.norm2 = nn.LayerNorm(quarter_dim)
|
||||
self.fc5 = nn.Linear(quarter_dim, quarter_dim)
|
||||
self.fc6 = nn.Linear(quarter_dim, 1)
|
||||
# set fc6 weights to near zero
|
||||
self.fc6.weight.data.normal_(mean=0.0, std=0.0001)
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
# x = (batch_size, 77, 768)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
x = self.norm1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc3(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc4(x)
|
||||
x = self.norm2(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc5(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc6(x)
|
||||
x = torch.sigmoid(x)
|
||||
return x
|
||||
|
||||
|
||||
class ZipperModule(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_size,
|
||||
in_tokens,
|
||||
out_size,
|
||||
out_tokens,
|
||||
hidden_size,
|
||||
hidden_tokens,
|
||||
use_residual=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_size = in_size
|
||||
self.in_tokens = in_tokens
|
||||
self.out_size = out_size
|
||||
self.out_tokens = out_tokens
|
||||
self.hidden_size = hidden_size
|
||||
self.hidden_tokens = hidden_tokens
|
||||
self.use_residual = use_residual
|
||||
|
||||
self.act_fn = nn.GELU()
|
||||
self.layernorm = nn.LayerNorm(self.in_size)
|
||||
|
||||
self.conv1 = nn.Conv1d(self.in_tokens, self.hidden_tokens, 1)
|
||||
# act
|
||||
self.fc1 = nn.Linear(self.in_size, self.hidden_size)
|
||||
# act
|
||||
self.conv2 = nn.Conv1d(self.hidden_tokens, self.out_tokens, 1)
|
||||
# act
|
||||
self.fc2 = nn.Linear(self.hidden_size, self.out_size)
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.conv1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.conv2(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class ZipperResampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_size,
|
||||
in_tokens,
|
||||
out_size,
|
||||
out_tokens,
|
||||
hidden_size,
|
||||
hidden_tokens,
|
||||
num_blocks=1,
|
||||
is_conv_input=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.is_conv_input = is_conv_input
|
||||
|
||||
module_list = []
|
||||
for i in range(num_blocks):
|
||||
|
||||
this_in_size = in_size
|
||||
this_in_tokens = in_tokens
|
||||
this_out_size = out_size
|
||||
this_out_tokens = out_tokens
|
||||
this_hidden_size = hidden_size
|
||||
this_hidden_tokens = hidden_tokens
|
||||
use_residual = False
|
||||
|
||||
# maintain middle sizes as hidden_size
|
||||
if i == 0: # first block
|
||||
this_in_size = in_size
|
||||
this_in_tokens = in_tokens
|
||||
if num_blocks == 1:
|
||||
this_out_size = out_size
|
||||
this_out_tokens = out_tokens
|
||||
else:
|
||||
this_out_size = hidden_size
|
||||
this_out_tokens = hidden_tokens
|
||||
elif i == num_blocks - 1: # last block
|
||||
this_out_size = out_size
|
||||
this_out_tokens = out_tokens
|
||||
if num_blocks == 1:
|
||||
this_in_size = in_size
|
||||
this_in_tokens = in_tokens
|
||||
else:
|
||||
this_in_size = hidden_size
|
||||
this_in_tokens = hidden_tokens
|
||||
else: # middle blocks
|
||||
this_out_size = hidden_size
|
||||
this_out_tokens = hidden_tokens
|
||||
this_in_size = hidden_size
|
||||
this_in_tokens = hidden_tokens
|
||||
use_residual = True
|
||||
|
||||
module_list.append(ZipperModule(
|
||||
in_size=this_in_size,
|
||||
in_tokens=this_in_tokens,
|
||||
out_size=this_out_size,
|
||||
out_tokens=this_out_tokens,
|
||||
hidden_size=this_hidden_size,
|
||||
hidden_tokens=this_hidden_tokens,
|
||||
use_residual=use_residual
|
||||
))
|
||||
|
||||
self.blocks = nn.ModuleList(module_list)
|
||||
|
||||
self.ctx_alpha = ContextualAlphaMask(
|
||||
dim=out_size,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
if self.is_conv_input:
|
||||
# flatten
|
||||
x = x.view(x.size(0), x.size(1), -1)
|
||||
# rearrange to (batch, tokens, size)
|
||||
x = x.permute(0, 2, 1)
|
||||
|
||||
for block in self.blocks:
|
||||
x = block(x)
|
||||
alpha = self.ctx_alpha(x)
|
||||
return x * alpha
|
||||
@@ -4,6 +4,7 @@ from collections import OrderedDict
|
||||
from typing import Optional, Union, List, Type, TYPE_CHECKING, Dict, Any, Literal
|
||||
|
||||
import torch
|
||||
from optimum.quanto import QTensor
|
||||
from torch import nn
|
||||
import weakref
|
||||
|
||||
@@ -13,14 +14,16 @@ from toolkit.config_modules import NetworkConfig
|
||||
from toolkit.lorm import extract_conv, extract_linear, count_parameters
|
||||
from toolkit.metadata import add_model_hash_to_meta
|
||||
from toolkit.paths import KEYMAPS_ROOT
|
||||
from toolkit.saving import get_lora_keymap_from_model_keymap
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lycoris_special import LycorisSpecialNetwork, LoConSpecialModule
|
||||
from toolkit.lora_special import LoRASpecialNetwork, LoRAModule
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.models.DoRA import DoRAModule
|
||||
|
||||
Network = Union['LycorisSpecialNetwork', 'LoRASpecialNetwork']
|
||||
Module = Union['LoConSpecialModule', 'LoRAModule']
|
||||
Module = Union['LoConSpecialModule', 'LoRAModule', 'DoRAModule']
|
||||
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
@@ -50,8 +53,14 @@ def broadcast_and_multiply(tensor, multiplier):
|
||||
for _ in range(num_extra_dims):
|
||||
multiplier = multiplier.unsqueeze(-1)
|
||||
|
||||
# Multiplying the broadcasted tensor with the output tensor
|
||||
result = tensor * multiplier
|
||||
try:
|
||||
# Multiplying the broadcasted tensor with the output tensor
|
||||
result = tensor * multiplier
|
||||
except RuntimeError as e:
|
||||
print(e)
|
||||
print(tensor.size())
|
||||
print(multiplier.size())
|
||||
raise e
|
||||
|
||||
return result
|
||||
|
||||
@@ -196,7 +205,6 @@ class ToolkitModuleMixin:
|
||||
|
||||
return lx * scale
|
||||
|
||||
|
||||
def lorm_forward(self: Network, x, *args, **kwargs):
|
||||
network: Network = self.network_ref()
|
||||
if not network.is_active:
|
||||
@@ -246,8 +254,17 @@ class ToolkitModuleMixin:
|
||||
# network is not active, avoid doing anything
|
||||
return self.org_forward(x, *args, **kwargs)
|
||||
|
||||
# if self.__class__.__name__ == "DoRAModule":
|
||||
# # return dora forward
|
||||
# return self.dora_forward(x, *args, **kwargs)
|
||||
|
||||
org_forwarded = self.org_forward(x, *args, **kwargs)
|
||||
lora_output = self._call_forward(x)
|
||||
|
||||
if isinstance(x, QTensor):
|
||||
x = x.dequantize()
|
||||
# always cast to float32
|
||||
lora_input = x.to(self.lora_down.weight.dtype)
|
||||
lora_output = self._call_forward(lora_input)
|
||||
multiplier = self.network_ref().torch_multiplier
|
||||
|
||||
lora_output_batch_size = lora_output.size(0)
|
||||
@@ -257,7 +274,34 @@ class ToolkitModuleMixin:
|
||||
# todo check if this is correct, do we just concat when doing cfg?
|
||||
multiplier = multiplier.repeat_interleave(num_interleaves)
|
||||
|
||||
x = org_forwarded + broadcast_and_multiply(lora_output, multiplier)
|
||||
scaled_lora_output = broadcast_and_multiply(lora_output, multiplier)
|
||||
scaled_lora_output = scaled_lora_output.to(org_forwarded.dtype)
|
||||
|
||||
if self.__class__.__name__ == "DoRAModule":
|
||||
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L417
|
||||
# x = dropout(x)
|
||||
# todo this wont match the dropout applied to the lora
|
||||
if isinstance(self.dropout, nn.Dropout) or isinstance(self.dropout, nn.Identity):
|
||||
lx = self.dropout(x)
|
||||
# normal dropout
|
||||
elif self.dropout is not None and self.training:
|
||||
lx = torch.nn.functional.dropout(x, p=self.dropout)
|
||||
else:
|
||||
lx = x
|
||||
lora_weight = self.lora_up.weight @ self.lora_down.weight
|
||||
# scale it here
|
||||
# todo handle our batch split scalers for slider training. For now take the mean of them
|
||||
scale = multiplier.mean()
|
||||
scaled_lora_weight = lora_weight * scale
|
||||
scaled_lora_output = scaled_lora_output + self.apply_dora(lx, scaled_lora_weight).to(org_forwarded.dtype)
|
||||
|
||||
try:
|
||||
x = org_forwarded + scaled_lora_output
|
||||
except RuntimeError as e:
|
||||
print(e)
|
||||
print(org_forwarded.size())
|
||||
print(scaled_lora_output.size())
|
||||
raise e
|
||||
return x
|
||||
|
||||
def enable_gradient_checkpointing(self: Module):
|
||||
@@ -275,14 +319,26 @@ class ToolkitModuleMixin:
|
||||
|
||||
@torch.no_grad()
|
||||
def merge_in(self: Module, merge_weight=1.0):
|
||||
if not self.can_merge_in:
|
||||
return
|
||||
# get up/down weight
|
||||
up_weight = self.lora_up.weight.clone().float()
|
||||
down_weight = self.lora_down.weight.clone().float()
|
||||
|
||||
# extract weight from org_module
|
||||
org_sd = self.org_module[0].state_dict()
|
||||
orig_dtype = org_sd["weight"].dtype
|
||||
weight = org_sd["weight"].float()
|
||||
# todo find a way to merge in weights when doing quantized model
|
||||
if 'weight._data' in org_sd:
|
||||
# quantized weight
|
||||
return
|
||||
|
||||
weight_key = "weight"
|
||||
if 'weight._data' in org_sd:
|
||||
# quantized weight
|
||||
weight_key = "weight._data"
|
||||
|
||||
orig_dtype = org_sd[weight_key].dtype
|
||||
weight = org_sd[weight_key].float()
|
||||
|
||||
multiplier = merge_weight
|
||||
scale = self.scale
|
||||
@@ -309,7 +365,7 @@ class ToolkitModuleMixin:
|
||||
weight = weight + multiplier * conved * scale
|
||||
|
||||
# set weight to org_module
|
||||
org_sd["weight"] = weight.to(orig_dtype)
|
||||
org_sd[weight_key] = weight.to(orig_dtype)
|
||||
self.org_module[0].load_state_dict(org_sd)
|
||||
|
||||
def setup_lorm(self: Module, state_dict: Optional[Dict[str, Any]] = None):
|
||||
@@ -338,6 +394,8 @@ class ToolkitNetworkMixin:
|
||||
train_unet: Optional[bool] = True,
|
||||
is_sdxl=False,
|
||||
is_v2=False,
|
||||
is_ssd=False,
|
||||
is_vega=False,
|
||||
network_config: Optional[NetworkConfig] = None,
|
||||
is_lorm=False,
|
||||
**kwargs
|
||||
@@ -348,7 +406,10 @@ class ToolkitNetworkMixin:
|
||||
self._multiplier: float = 1.0
|
||||
self.is_active: bool = False
|
||||
self.is_sdxl = is_sdxl
|
||||
self.is_ssd = is_ssd
|
||||
self.is_vega = is_vega
|
||||
self.is_v2 = is_v2
|
||||
self.is_v1 = not is_v2 and not is_sdxl and not is_ssd and not is_vega
|
||||
self.is_merged_in = False
|
||||
self.is_lorm = is_lorm
|
||||
self.network_config: NetworkConfig = network_config
|
||||
@@ -356,15 +417,32 @@ class ToolkitNetworkMixin:
|
||||
self.lorm_train_mode: Literal['local', None] = None
|
||||
self.can_merge_in = not is_lorm
|
||||
|
||||
def get_keymap(self: Network):
|
||||
if self.is_sdxl:
|
||||
def get_keymap(self: Network, force_weight_mapping=False):
|
||||
use_weight_mapping = False
|
||||
|
||||
if self.is_ssd:
|
||||
keymap_tail = 'ssd'
|
||||
use_weight_mapping = True
|
||||
elif self.is_vega:
|
||||
keymap_tail = 'vega'
|
||||
use_weight_mapping = True
|
||||
elif self.is_sdxl:
|
||||
keymap_tail = 'sdxl'
|
||||
elif self.is_v2:
|
||||
keymap_tail = 'sd2'
|
||||
else:
|
||||
keymap_tail = 'sd1'
|
||||
# todo double check this
|
||||
# use_weight_mapping = True
|
||||
|
||||
if force_weight_mapping:
|
||||
use_weight_mapping = True
|
||||
|
||||
# load keymap
|
||||
keymap_name = f"stable_diffusion_locon_{keymap_tail}.json"
|
||||
if use_weight_mapping:
|
||||
keymap_name = f"stable_diffusion_{keymap_tail}.json"
|
||||
|
||||
keymap_path = os.path.join(KEYMAPS_ROOT, keymap_name)
|
||||
|
||||
keymap = None
|
||||
@@ -373,6 +451,27 @@ class ToolkitNetworkMixin:
|
||||
with open(keymap_path, 'r') as f:
|
||||
keymap = json.load(f)['ldm_diffusers_keymap']
|
||||
|
||||
if use_weight_mapping and keymap is not None:
|
||||
# get keymap from weights
|
||||
keymap = get_lora_keymap_from_model_keymap(keymap)
|
||||
|
||||
# upgrade keymaps for DoRA
|
||||
if self.network_type.lower() == 'dora':
|
||||
if keymap is not None:
|
||||
new_keymap = {}
|
||||
for ldm_key, diffusers_key in keymap.items():
|
||||
ldm_key = ldm_key.replace('.alpha', '.magnitude')
|
||||
# ldm_key = ldm_key.replace('.lora_down.weight', '.lora_down')
|
||||
# ldm_key = ldm_key.replace('.lora_up.weight', '.lora_up')
|
||||
|
||||
diffusers_key = diffusers_key.replace('.alpha', '.magnitude')
|
||||
# diffusers_key = diffusers_key.replace('.lora_down.weight', '.lora_down')
|
||||
# diffusers_key = diffusers_key.replace('.lora_up.weight', '.lora_up')
|
||||
|
||||
new_keymap[ldm_key] = diffusers_key
|
||||
|
||||
keymap = new_keymap
|
||||
|
||||
return keymap
|
||||
|
||||
def save_weights(
|
||||
@@ -400,6 +499,7 @@ class ToolkitNetworkMixin:
|
||||
v = v.detach().clone().to("cpu").to(dtype)
|
||||
save_key = save_keymap[key] if key in save_keymap else key
|
||||
save_dict[save_key] = v
|
||||
del state_dict[key]
|
||||
|
||||
if extra_state_dict is not None:
|
||||
# add extra items to state dict
|
||||
@@ -408,6 +508,24 @@ class ToolkitNetworkMixin:
|
||||
v = v.detach().clone().to("cpu").to(dtype)
|
||||
save_dict[key] = v
|
||||
|
||||
if self.peft_format:
|
||||
# lora_down = lora_A
|
||||
# lora_up = lora_B
|
||||
# no alpha
|
||||
|
||||
new_save_dict = {}
|
||||
for key, value in save_dict.items():
|
||||
if key.endswith('.alpha'):
|
||||
continue
|
||||
new_key = key
|
||||
new_key = new_key.replace('lora_down', 'lora_A')
|
||||
new_key = new_key.replace('lora_up', 'lora_B')
|
||||
# replace all $$ with .
|
||||
new_key = new_key.replace('$$', '.')
|
||||
new_save_dict[new_key] = value
|
||||
|
||||
save_dict = new_save_dict
|
||||
|
||||
if metadata is None:
|
||||
metadata = OrderedDict()
|
||||
metadata = add_model_hash_to_meta(state_dict, metadata)
|
||||
@@ -417,21 +535,42 @@ class ToolkitNetworkMixin:
|
||||
else:
|
||||
torch.save(save_dict, file)
|
||||
|
||||
def load_weights(self: Network, file):
|
||||
def load_weights(self: Network, file, force_weight_mapping=False):
|
||||
# allows us to save and load to and from ldm weights
|
||||
keymap = self.get_keymap()
|
||||
keymap = self.get_keymap(force_weight_mapping)
|
||||
keymap = {} if keymap is None else keymap
|
||||
|
||||
if os.path.splitext(file)[1] == ".safetensors":
|
||||
from safetensors.torch import load_file
|
||||
if isinstance(file, str):
|
||||
if os.path.splitext(file)[1] == ".safetensors":
|
||||
from safetensors.torch import load_file
|
||||
|
||||
weights_sd = load_file(file)
|
||||
weights_sd = load_file(file)
|
||||
else:
|
||||
weights_sd = torch.load(file, map_location="cpu")
|
||||
else:
|
||||
weights_sd = torch.load(file, map_location="cpu")
|
||||
# probably a state dict
|
||||
weights_sd = file
|
||||
|
||||
load_sd = OrderedDict()
|
||||
for key, value in weights_sd.items():
|
||||
load_key = keymap[key] if key in keymap else key
|
||||
# replace old double __ with single _
|
||||
if self.is_pixart:
|
||||
load_key = load_key.replace('__', '_')
|
||||
|
||||
if self.peft_format:
|
||||
# lora_down = lora_A
|
||||
# lora_up = lora_B
|
||||
# no alpha
|
||||
if load_key.endswith('.alpha'):
|
||||
continue
|
||||
load_key = load_key.replace('lora_A', 'lora_down')
|
||||
load_key = load_key.replace('lora_B', 'lora_up')
|
||||
# replace all . with $$
|
||||
load_key = load_key.replace('.', '$$')
|
||||
load_key = load_key.replace('$$lora_down$$', '.lora_down.')
|
||||
load_key = load_key.replace('$$lora_up$$', '.lora_up.')
|
||||
|
||||
load_sd[load_key] = value
|
||||
|
||||
# extract extra items from state dict
|
||||
@@ -445,6 +584,12 @@ class ToolkitNetworkMixin:
|
||||
for key in to_delete:
|
||||
del load_sd[key]
|
||||
|
||||
print(f"Missing keys: {to_delete}")
|
||||
if len(to_delete) > 0 and self.is_v1 and not force_weight_mapping and not (
|
||||
len(to_delete) == 1 and 'emb_params' in to_delete):
|
||||
print(" Attempting to load with forced keymap")
|
||||
return self.load_weights(file, force_weight_mapping=True)
|
||||
|
||||
info = self.load_state_dict(load_sd, False)
|
||||
if len(extra_dict.keys()) == 0:
|
||||
extra_dict = None
|
||||
@@ -527,6 +672,8 @@ class ToolkitNetworkMixin:
|
||||
self._update_checkpointing()
|
||||
|
||||
def merge_in(self, merge_weight=1.0):
|
||||
if self.network_type.lower() == 'dora':
|
||||
return
|
||||
self.is_merged_in = True
|
||||
for module in self.get_all_modules():
|
||||
module.merge_in(merge_weight)
|
||||
@@ -563,4 +710,3 @@ class ToolkitNetworkMixin:
|
||||
params_reduced += (num_orig_module_params - num_lorem_params)
|
||||
|
||||
return params_reduced
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import torch
|
||||
from transformers import Adafactor
|
||||
from transformers import Adafactor, AdamW
|
||||
|
||||
|
||||
def get_optimizer(
|
||||
@@ -20,12 +20,12 @@ def get_optimizer(
|
||||
# dadaptation uses different lr that is values of 0.1 to 1.0. default to 1.0
|
||||
use_lr = 1.0
|
||||
if lower_type.endswith('lion'):
|
||||
optimizer = dadaptation.DAdaptLion(params, lr=use_lr, **optimizer_params)
|
||||
optimizer = dadaptation.DAdaptLion(params, eps=1e-6, lr=use_lr, **optimizer_params)
|
||||
elif lower_type.endswith('adam'):
|
||||
optimizer = dadaptation.DAdaptLion(params, lr=use_lr, **optimizer_params)
|
||||
optimizer = dadaptation.DAdaptLion(params, eps=1e-6, lr=use_lr, **optimizer_params)
|
||||
elif lower_type == 'dadaptation':
|
||||
# backwards compatibility
|
||||
optimizer = dadaptation.DAdaptAdam(params, lr=use_lr, **optimizer_params)
|
||||
optimizer = dadaptation.DAdaptAdam(params, eps=1e-6, lr=use_lr, **optimizer_params)
|
||||
# warn user that dadaptation is deprecated
|
||||
print("WARNING: Dadaptation optimizer type has been changed to DadaptationAdam. Please update your config.")
|
||||
elif lower_type.startswith("prodigy"):
|
||||
@@ -40,22 +40,22 @@ def get_optimizer(
|
||||
print(f"Using lr {use_lr}")
|
||||
# let net be the neural network you want to train
|
||||
# you can choose weight decay value based on your problem, 0 by default
|
||||
optimizer = Prodigy(params, lr=use_lr, **optimizer_params)
|
||||
optimizer = Prodigy(params, lr=use_lr, eps=1e-6, **optimizer_params)
|
||||
elif lower_type.endswith("8bit"):
|
||||
import bitsandbytes
|
||||
|
||||
if lower_type == "adam8bit":
|
||||
return bitsandbytes.optim.Adam8bit(params, lr=learning_rate, **optimizer_params)
|
||||
return bitsandbytes.optim.Adam8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "adamw8bit":
|
||||
return bitsandbytes.optim.AdamW8bit(params, lr=learning_rate, **optimizer_params)
|
||||
return bitsandbytes.optim.AdamW8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "lion8bit":
|
||||
return bitsandbytes.optim.Lion8bit(params, lr=learning_rate, **optimizer_params)
|
||||
else:
|
||||
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
||||
elif lower_type == 'adam':
|
||||
optimizer = torch.optim.Adam(params, lr=float(learning_rate), **optimizer_params)
|
||||
optimizer = torch.optim.Adam(params, lr=float(learning_rate), eps=1e-6, **optimizer_params)
|
||||
elif lower_type == 'adamw':
|
||||
optimizer = torch.optim.AdamW(params, lr=float(learning_rate), **optimizer_params)
|
||||
optimizer = torch.optim.AdamW(params, lr=float(learning_rate), eps=1e-6, **optimizer_params)
|
||||
elif lower_type == 'lion':
|
||||
try:
|
||||
from lion_pytorch import Lion
|
||||
@@ -63,9 +63,18 @@ def get_optimizer(
|
||||
except ImportError:
|
||||
raise ImportError("Please install lion_pytorch to use Lion optimizer -> pip install lion-pytorch")
|
||||
elif lower_type == 'adagrad':
|
||||
optimizer = torch.optim.Adagrad(params, lr=float(learning_rate), **optimizer_params)
|
||||
optimizer = torch.optim.Adagrad(params, lr=float(learning_rate), eps=1e-6, **optimizer_params)
|
||||
elif lower_type == 'adafactor':
|
||||
optimizer = Adafactor(params, lr=float(learning_rate), **optimizer_params)
|
||||
# hack in stochastic rounding
|
||||
if 'relative_step' not in optimizer_params:
|
||||
optimizer_params['relative_step'] = False
|
||||
if 'scale_parameter' not in optimizer_params:
|
||||
optimizer_params['scale_parameter'] = False
|
||||
if 'warmup_init' not in optimizer_params:
|
||||
optimizer_params['warmup_init'] = False
|
||||
optimizer = Adafactor(params, lr=float(learning_rate), eps=1e-6, **optimizer_params)
|
||||
from toolkit.util.adafactor_stochastic_rounding import step_adafactor
|
||||
optimizer.step = step_adafactor.__get__(optimizer, Adafactor)
|
||||
else:
|
||||
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
||||
return optimizer
|
||||
|
||||
144
toolkit/photomaker.py
Normal file
144
toolkit/photomaker.py
Normal file
@@ -0,0 +1,144 @@
|
||||
# Merge image encoder and fuse module to create an ID Encoder
|
||||
# send multiple ID images, we can directly obtain the updated text encoder containing a stacked ID embedding
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers.models.clip.modeling_clip import CLIPVisionModelWithProjection
|
||||
from transformers.models.clip.configuration_clip import CLIPVisionConfig
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
VISION_CONFIG_DICT = {
|
||||
"hidden_size": 1024,
|
||||
"intermediate_size": 4096,
|
||||
"num_attention_heads": 16,
|
||||
"num_hidden_layers": 24,
|
||||
"patch_size": 14,
|
||||
"projection_dim": 768
|
||||
}
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = nn.LayerNorm(in_dim)
|
||||
self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
self.use_residual = use_residual
|
||||
self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class FuseModule(nn.Module):
|
||||
def __init__(self, embed_dim):
|
||||
super().__init__()
|
||||
self.mlp1 = MLP(embed_dim * 2, embed_dim, embed_dim, use_residual=False)
|
||||
self.mlp2 = MLP(embed_dim, embed_dim, embed_dim, use_residual=True)
|
||||
self.layer_norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
def fuse_fn(self, prompt_embeds, id_embeds):
|
||||
stacked_id_embeds = torch.cat([prompt_embeds, id_embeds], dim=-1)
|
||||
stacked_id_embeds = self.mlp1(stacked_id_embeds) + prompt_embeds
|
||||
stacked_id_embeds = self.mlp2(stacked_id_embeds)
|
||||
stacked_id_embeds = self.layer_norm(stacked_id_embeds)
|
||||
return stacked_id_embeds
|
||||
|
||||
def forward(
|
||||
self,
|
||||
prompt_embeds,
|
||||
id_embeds,
|
||||
class_tokens_mask,
|
||||
) -> torch.Tensor:
|
||||
# id_embeds shape: [b, max_num_inputs, 1, 2048]
|
||||
id_embeds = id_embeds.to(prompt_embeds.dtype)
|
||||
num_inputs = class_tokens_mask.sum().unsqueeze(0) # TODO: check for training case
|
||||
batch_size, max_num_inputs = id_embeds.shape[:2]
|
||||
# seq_length: 77
|
||||
seq_length = prompt_embeds.shape[1]
|
||||
# flat_id_embeds shape: [b*max_num_inputs, 1, 2048]
|
||||
flat_id_embeds = id_embeds.view(
|
||||
-1, id_embeds.shape[-2], id_embeds.shape[-1]
|
||||
)
|
||||
# valid_id_mask [b*max_num_inputs]
|
||||
valid_id_mask = (
|
||||
torch.arange(max_num_inputs, device=flat_id_embeds.device)[None, :]
|
||||
< num_inputs[:, None]
|
||||
)
|
||||
valid_id_embeds = flat_id_embeds[valid_id_mask.flatten()]
|
||||
|
||||
prompt_embeds = prompt_embeds.view(-1, prompt_embeds.shape[-1])
|
||||
class_tokens_mask = class_tokens_mask.view(-1)
|
||||
valid_id_embeds = valid_id_embeds.view(-1, valid_id_embeds.shape[-1])
|
||||
# slice out the image token embeddings
|
||||
image_token_embeds = prompt_embeds[class_tokens_mask]
|
||||
stacked_id_embeds = self.fuse_fn(image_token_embeds, valid_id_embeds)
|
||||
assert class_tokens_mask.sum() == stacked_id_embeds.shape[0], f"{class_tokens_mask.sum()} != {stacked_id_embeds.shape[0]}"
|
||||
prompt_embeds.masked_scatter_(class_tokens_mask[:, None], stacked_id_embeds.to(prompt_embeds.dtype))
|
||||
updated_prompt_embeds = prompt_embeds.view(batch_size, seq_length, -1)
|
||||
return updated_prompt_embeds
|
||||
|
||||
class PhotoMakerIDEncoder(CLIPVisionModelWithProjection):
|
||||
def __init__(self, config=None, *model_args, **model_kwargs):
|
||||
if config is None:
|
||||
config = CLIPVisionConfig(**VISION_CONFIG_DICT)
|
||||
super().__init__(config, *model_args, **model_kwargs)
|
||||
self.visual_projection_2 = nn.Linear(1024, 1280, bias=False)
|
||||
self.fuse_module = FuseModule(2048)
|
||||
|
||||
def forward(self, id_pixel_values, prompt_embeds, class_tokens_mask):
|
||||
b, num_inputs, c, h, w = id_pixel_values.shape
|
||||
id_pixel_values = id_pixel_values.view(b * num_inputs, c, h, w)
|
||||
|
||||
shared_id_embeds = self.vision_model(id_pixel_values)[1]
|
||||
id_embeds = self.visual_projection(shared_id_embeds)
|
||||
id_embeds_2 = self.visual_projection_2(shared_id_embeds)
|
||||
|
||||
id_embeds = id_embeds.view(b, num_inputs, 1, -1)
|
||||
id_embeds_2 = id_embeds_2.view(b, num_inputs, 1, -1)
|
||||
|
||||
id_embeds = torch.cat((id_embeds, id_embeds_2), dim=-1)
|
||||
updated_prompt_embeds = self.fuse_module(
|
||||
prompt_embeds, id_embeds, class_tokens_mask)
|
||||
|
||||
return updated_prompt_embeds
|
||||
|
||||
|
||||
class PhotoMakerCLIPEncoder(CLIPVisionModelWithProjection):
|
||||
def __init__(self, config=None, *model_args, **model_kwargs):
|
||||
if config is None:
|
||||
config = CLIPVisionConfig(**VISION_CONFIG_DICT)
|
||||
super().__init__(config, *model_args, **model_kwargs)
|
||||
self.visual_projection_2 = nn.Linear(1024, 1280, bias=False)
|
||||
|
||||
def forward(self, id_pixel_values, do_projection2=True, output_full=False):
|
||||
b, num_inputs, c, h, w = id_pixel_values.shape
|
||||
id_pixel_values = id_pixel_values.view(b * num_inputs, c, h, w)
|
||||
# last_hidden_state, 1, 257, 1024
|
||||
vision_output = self.vision_model(id_pixel_values, output_hidden_states=True)
|
||||
shared_id_embeds = vision_output[1]
|
||||
id_embeds = self.visual_projection(shared_id_embeds)
|
||||
|
||||
id_embeds = id_embeds.view(b, num_inputs, 1, -1)
|
||||
|
||||
if do_projection2:
|
||||
id_embeds_2 = self.visual_projection_2(shared_id_embeds)
|
||||
id_embeds_2 = id_embeds_2.view(b, num_inputs, 1, -1)
|
||||
id_embeds = torch.cat((id_embeds, id_embeds_2), dim=-1)
|
||||
|
||||
if output_full:
|
||||
return id_embeds, vision_output
|
||||
return id_embeds
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
PhotoMakerIDEncoder()
|
||||
491
toolkit/photomaker_pipeline.py
Normal file
491
toolkit/photomaker_pipeline.py
Normal file
@@ -0,0 +1,491 @@
|
||||
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
import PIL
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
from torchvision import transforms as T
|
||||
|
||||
from safetensors import safe_open
|
||||
from huggingface_hub.utils import validate_hf_hub_args
|
||||
from transformers import CLIPImageProcessor, CLIPTokenizer
|
||||
from diffusers import StableDiffusionXLPipeline
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput
|
||||
from diffusers.utils import (
|
||||
_get_model_file,
|
||||
is_transformers_available,
|
||||
logging,
|
||||
)
|
||||
|
||||
from .photomaker import PhotoMakerIDEncoder
|
||||
|
||||
PipelineImageInput = Union[
|
||||
PIL.Image.Image,
|
||||
torch.FloatTensor,
|
||||
List[PIL.Image.Image],
|
||||
List[torch.FloatTensor],
|
||||
]
|
||||
|
||||
|
||||
class PhotoMakerStableDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
@validate_hf_hub_args
|
||||
def load_photomaker_adapter(
|
||||
self,
|
||||
pretrained_model_name_or_path_or_dict: Union[str, Dict[str, torch.Tensor]],
|
||||
weight_name: str,
|
||||
subfolder: str = '',
|
||||
trigger_word: str = 'img',
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Parameters:
|
||||
pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`):
|
||||
Can be either:
|
||||
|
||||
- A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on
|
||||
the Hub.
|
||||
- A path to a *directory* (for example `./my_model_directory`) containing the model weights saved
|
||||
with [`ModelMixin.save_pretrained`].
|
||||
- A [torch state
|
||||
dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict).
|
||||
|
||||
weight_name (`str`):
|
||||
The weight name NOT the path to the weight.
|
||||
|
||||
subfolder (`str`, defaults to `""`):
|
||||
The subfolder location of a model file within a larger model repository on the Hub or locally.
|
||||
|
||||
trigger_word (`str`, *optional*, defaults to `"img"`):
|
||||
The trigger word is used to identify the position of class word in the text prompt,
|
||||
and it is recommended not to set it as a common word.
|
||||
This trigger word must be placed after the class word when used, otherwise, it will affect the performance of the personalized generation.
|
||||
"""
|
||||
|
||||
# Load the main state dict first.
|
||||
cache_dir = kwargs.pop("cache_dir", None)
|
||||
force_download = kwargs.pop("force_download", False)
|
||||
resume_download = kwargs.pop("resume_download", False)
|
||||
proxies = kwargs.pop("proxies", None)
|
||||
local_files_only = kwargs.pop("local_files_only", None)
|
||||
token = kwargs.pop("token", None)
|
||||
revision = kwargs.pop("revision", None)
|
||||
|
||||
user_agent = {
|
||||
"file_type": "attn_procs_weights",
|
||||
"framework": "pytorch",
|
||||
}
|
||||
|
||||
if not isinstance(pretrained_model_name_or_path_or_dict, dict):
|
||||
model_file = _get_model_file(
|
||||
pretrained_model_name_or_path_or_dict,
|
||||
weights_name=weight_name,
|
||||
cache_dir=cache_dir,
|
||||
force_download=force_download,
|
||||
resume_download=resume_download,
|
||||
proxies=proxies,
|
||||
local_files_only=local_files_only,
|
||||
token=token,
|
||||
revision=revision,
|
||||
subfolder=subfolder,
|
||||
user_agent=user_agent,
|
||||
)
|
||||
if weight_name.endswith(".safetensors"):
|
||||
state_dict = {"id_encoder": {}, "lora_weights": {}}
|
||||
with safe_open(model_file, framework="pt", device="cpu") as f:
|
||||
for key in f.keys():
|
||||
if key.startswith("id_encoder."):
|
||||
state_dict["id_encoder"][key.replace("id_encoder.", "")] = f.get_tensor(key)
|
||||
elif key.startswith("lora_weights."):
|
||||
state_dict["lora_weights"][key.replace("lora_weights.", "")] = f.get_tensor(key)
|
||||
else:
|
||||
state_dict = torch.load(model_file, map_location="cpu")
|
||||
else:
|
||||
state_dict = pretrained_model_name_or_path_or_dict
|
||||
|
||||
keys = list(state_dict.keys())
|
||||
if keys != ["id_encoder", "lora_weights"]:
|
||||
raise ValueError("Required keys are (`id_encoder` and `lora_weights`) missing from the state dict.")
|
||||
|
||||
self.trigger_word = trigger_word
|
||||
# load finetuned CLIP image encoder and fuse module here if it has not been registered to the pipeline yet
|
||||
print(f"Loading PhotoMaker components [1] id_encoder from [{pretrained_model_name_or_path_or_dict}]...")
|
||||
id_encoder = PhotoMakerIDEncoder()
|
||||
id_encoder.load_state_dict(state_dict["id_encoder"], strict=True)
|
||||
id_encoder = id_encoder.to(self.device, dtype=self.unet.dtype)
|
||||
self.id_encoder = id_encoder
|
||||
self.id_image_processor = CLIPImageProcessor()
|
||||
|
||||
# load lora into models
|
||||
print(f"Loading PhotoMaker components [2] lora_weights from [{pretrained_model_name_or_path_or_dict}]")
|
||||
self.load_lora_weights(state_dict["lora_weights"], adapter_name="photomaker")
|
||||
|
||||
# Add trigger word token
|
||||
if self.tokenizer is not None:
|
||||
self.tokenizer.add_tokens([self.trigger_word], special_tokens=True)
|
||||
|
||||
self.tokenizer_2.add_tokens([self.trigger_word], special_tokens=True)
|
||||
|
||||
def encode_prompt_with_trigger_word(
|
||||
self,
|
||||
prompt: str,
|
||||
prompt_2: Optional[str] = None,
|
||||
num_id_images: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
class_tokens_mask: Optional[torch.LongTensor] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
# Find the token id of the trigger word
|
||||
image_token_id = self.tokenizer_2.convert_tokens_to_ids(self.trigger_word)
|
||||
|
||||
# Define tokenizers and text encoders
|
||||
tokenizers = [self.tokenizer, self.tokenizer_2] if self.tokenizer is not None else [self.tokenizer_2]
|
||||
text_encoders = (
|
||||
[self.text_encoder, self.text_encoder_2] if self.text_encoder is not None else [self.text_encoder_2]
|
||||
)
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_2 = prompt_2 or prompt
|
||||
prompt_embeds_list = []
|
||||
prompts = [prompt, prompt_2]
|
||||
for prompt, tokenizer, text_encoder in zip(prompts, tokenizers, text_encoders):
|
||||
input_ids = tokenizer.encode(prompt) # TODO: batch encode
|
||||
clean_index = 0
|
||||
clean_input_ids = []
|
||||
class_token_index = []
|
||||
# Find out the corrresponding class word token based on the newly added trigger word token
|
||||
for i, token_id in enumerate(input_ids):
|
||||
if token_id == image_token_id:
|
||||
class_token_index.append(clean_index - 1)
|
||||
else:
|
||||
clean_input_ids.append(token_id)
|
||||
clean_index += 1
|
||||
|
||||
if len(class_token_index) != 1:
|
||||
raise ValueError(
|
||||
f"PhotoMaker currently does not support multiple trigger words in a single prompt.\
|
||||
Trigger word: {self.trigger_word}, Prompt: {prompt}."
|
||||
)
|
||||
class_token_index = class_token_index[0]
|
||||
|
||||
# Expand the class word token and corresponding mask
|
||||
class_token = clean_input_ids[class_token_index]
|
||||
clean_input_ids = clean_input_ids[:class_token_index] + [class_token] * num_id_images + \
|
||||
clean_input_ids[class_token_index + 1:]
|
||||
|
||||
# Truncation or padding
|
||||
max_len = tokenizer.model_max_length
|
||||
if len(clean_input_ids) > max_len:
|
||||
clean_input_ids = clean_input_ids[:max_len]
|
||||
else:
|
||||
clean_input_ids = clean_input_ids + [tokenizer.pad_token_id] * (
|
||||
max_len - len(clean_input_ids)
|
||||
)
|
||||
|
||||
class_tokens_mask = [True if class_token_index <= i < class_token_index + num_id_images else False \
|
||||
for i in range(len(clean_input_ids))]
|
||||
|
||||
clean_input_ids = torch.tensor(clean_input_ids, dtype=torch.long).unsqueeze(0)
|
||||
class_tokens_mask = torch.tensor(class_tokens_mask, dtype=torch.bool).unsqueeze(0)
|
||||
|
||||
prompt_embeds = text_encoder(
|
||||
clean_input_ids.to(device),
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
# We are only ALWAYS interested in the pooled output of the final text encoder
|
||||
pooled_prompt_embeds = prompt_embeds[0]
|
||||
prompt_embeds = prompt_embeds.hidden_states[-2]
|
||||
prompt_embeds_list.append(prompt_embeds)
|
||||
|
||||
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder_2.dtype, device=device)
|
||||
class_tokens_mask = class_tokens_mask.to(device=device) # TODO: ignoring two-prompt case
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, class_tokens_mask
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
denoising_end: Optional[float] = None,
|
||||
guidance_scale: float = 5.0,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
guidance_rescale: float = 0.0,
|
||||
original_size: Optional[Tuple[int, int]] = None,
|
||||
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||
target_size: Optional[Tuple[int, int]] = None,
|
||||
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||
callback_steps: int = 1,
|
||||
# Added parameters (for PhotoMaker)
|
||||
input_id_images: PipelineImageInput = None,
|
||||
start_merge_step: int = 0, # TODO: change to `style_strength_ratio` in the future
|
||||
class_tokens_mask: Optional[torch.LongTensor] = None,
|
||||
prompt_embeds_text_only: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds_text_only: Optional[torch.FloatTensor] = None,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
Only the parameters introduced by PhotoMaker are discussed here.
|
||||
For explanations of the previous parameters in StableDiffusionXLPipeline, please refer to https://github.com/huggingface/diffusers/blob/v0.25.0/src/diffusers/pipelines/stable_diffusion_xl/pipeline_stable_diffusion_xl.py
|
||||
|
||||
Args:
|
||||
input_id_images (`PipelineImageInput`, *optional*):
|
||||
Input ID Image to work with PhotoMaker.
|
||||
class_tokens_mask (`torch.LongTensor`, *optional*):
|
||||
Pre-generated class token. When the `prompt_embeds` parameter is provided in advance, it is necessary to prepare the `class_tokens_mask` beforehand for marking out the position of class word.
|
||||
prompt_embeds_text_only (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds_text_only (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] or `tuple`:
|
||||
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] if `return_dict` is True, otherwise a
|
||||
`tuple`. When returning a tuple, the first element is a list with the generated images.
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
original_size = original_size or (height, width)
|
||||
target_size = target_size or (height, width)
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
callback_steps,
|
||||
negative_prompt,
|
||||
negative_prompt_2,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
)
|
||||
#
|
||||
if prompt_embeds is not None and class_tokens_mask is None:
|
||||
raise ValueError(
|
||||
"If `prompt_embeds` are provided, `class_tokens_mask` also have to be passed. Make sure to generate `class_tokens_mask` from the same tokenizer that was used to generate `prompt_embeds`."
|
||||
)
|
||||
# check the input id images
|
||||
if input_id_images is None:
|
||||
raise ValueError(
|
||||
"Provide `input_id_images`. Cannot leave `input_id_images` undefined for PhotoMaker pipeline."
|
||||
)
|
||||
if not isinstance(input_id_images, list):
|
||||
input_id_images = [input_id_images]
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
assert do_classifier_free_guidance
|
||||
|
||||
# 3. Encode input prompt
|
||||
num_id_images = len(input_id_images)
|
||||
|
||||
(
|
||||
prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
class_tokens_mask,
|
||||
) = self.encode_prompt_with_trigger_word(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
device=device,
|
||||
num_id_images=num_id_images,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
class_tokens_mask=class_tokens_mask,
|
||||
)
|
||||
|
||||
# 4. Encode input prompt without the trigger word for delayed conditioning
|
||||
prompt_text_only = prompt.replace(" " + self.trigger_word, "") # sensitive to white space
|
||||
(
|
||||
prompt_embeds_text_only,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds_text_only, # TODO: replace the pooled_prompt_embeds with text only prompt
|
||||
negative_pooled_prompt_embeds,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt_text_only,
|
||||
prompt_2=prompt_2,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||
negative_prompt=negative_prompt,
|
||||
negative_prompt_2=negative_prompt_2,
|
||||
prompt_embeds=prompt_embeds_text_only,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds_text_only,
|
||||
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
)
|
||||
|
||||
# 5. Prepare the input ID images
|
||||
dtype = next(self.id_encoder.parameters()).dtype
|
||||
if not isinstance(input_id_images[0], torch.Tensor):
|
||||
id_pixel_values = self.id_image_processor(input_id_images, return_tensors="pt").pixel_values
|
||||
|
||||
id_pixel_values = id_pixel_values.unsqueeze(0).to(device=device, dtype=dtype) # TODO: multiple prompts
|
||||
|
||||
# 6. Get the update text embedding with the stacked ID embedding
|
||||
prompt_embeds = self.id_encoder(id_pixel_values, prompt_embeds, class_tokens_mask)
|
||||
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
|
||||
pooled_prompt_embeds = pooled_prompt_embeds.repeat(1, num_images_per_prompt).view(
|
||||
bs_embed * num_images_per_prompt, -1
|
||||
)
|
||||
|
||||
# 7. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 8. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
# 10. Prepare added time ids & embeddings
|
||||
if self.text_encoder_2 is None:
|
||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||
else:
|
||||
text_encoder_projection_dim = self.text_encoder_2.config.projection_dim
|
||||
|
||||
add_time_ids = self._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 = torch.cat([add_time_ids, add_time_ids], dim=0)
|
||||
add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1)
|
||||
|
||||
# 11. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
latent_model_input = (
|
||||
torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
)
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
if i <= start_merge_step:
|
||||
current_prompt_embeds = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds_text_only], dim=0
|
||||
)
|
||||
add_text_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_prompt_embeds_text_only], dim=0)
|
||||
else:
|
||||
current_prompt_embeds = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds], dim=0
|
||||
)
|
||||
add_text_embeds = torch.cat([negative_pooled_prompt_embeds, pooled_prompt_embeds], dim=0)
|
||||
# predict the noise residual
|
||||
added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids}
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=current_prompt_embeds,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
if do_classifier_free_guidance and guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
callback(i, t, latents)
|
||||
|
||||
# make sure the VAE is in float32 mode, as it overflows in float16
|
||||
if self.vae.dtype == torch.float16 and self.vae.config.force_upcast:
|
||||
self.upcast_vae()
|
||||
latents = latents.to(next(iter(self.vae.post_quant_conv.parameters())).dtype)
|
||||
|
||||
if not output_type == "latent":
|
||||
image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0]
|
||||
else:
|
||||
image = latents
|
||||
return StableDiffusionXLPipelineOutput(images=image)
|
||||
|
||||
# apply watermark if available
|
||||
# if self.watermark is not None:
|
||||
# image = self.watermark.apply_watermark(image)
|
||||
|
||||
image = self.image_processor.postprocess(image, output_type=output_type)
|
||||
|
||||
# Offload last model to CPU
|
||||
if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None:
|
||||
self.final_offload_hook.offload()
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return StableDiffusionXLPipelineOutput(images=image)
|
||||
@@ -2,11 +2,14 @@ import importlib
|
||||
import inspect
|
||||
from typing import Union, List, Optional, Dict, Any, Tuple, Callable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionPipeline, LMSDiscreteScheduler
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionPipeline, LMSDiscreteScheduler, FluxPipeline
|
||||
from diffusers.pipelines.flux.pipeline_flux import calculate_shift, retrieve_timesteps
|
||||
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_k_diffusion import ModelWrapper
|
||||
from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput
|
||||
# from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_k_diffusion import ModelWrapper
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import rescale_noise_cfg
|
||||
from diffusers.utils import is_torch_xla_available
|
||||
from k_diffusion.external import CompVisVDenoiser, CompVisDenoiser
|
||||
@@ -43,13 +46,14 @@ class StableDiffusionKDiffusionXLPipeline(StableDiffusionXLPipeline):
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
self.sampler = None
|
||||
scheduler = LMSDiscreteScheduler.from_config(scheduler.config)
|
||||
model = ModelWrapper(unet, scheduler.alphas_cumprod)
|
||||
if scheduler.config.prediction_type == "v_prediction":
|
||||
self.k_diffusion_model = CompVisVDenoiser(model)
|
||||
else:
|
||||
self.k_diffusion_model = CompVisDenoiser(model)
|
||||
raise NotImplementedError("This pipeline is not implemented yet")
|
||||
# self.sampler = None
|
||||
# scheduler = LMSDiscreteScheduler.from_config(scheduler.config)
|
||||
# model = ModelWrapper(unet, scheduler.alphas_cumprod)
|
||||
# if scheduler.config.prediction_type == "v_prediction":
|
||||
# self.k_diffusion_model = CompVisVDenoiser(model)
|
||||
# else:
|
||||
# self.k_diffusion_model = CompVisDenoiser(model)
|
||||
|
||||
def set_scheduler(self, scheduler_type: str):
|
||||
library = importlib.import_module("k_diffusion")
|
||||
@@ -1201,3 +1205,213 @@ class StableDiffusionXLRefinerPipeline(StableDiffusionXLPipeline):
|
||||
|
||||
return StableDiffusionXLPipelineOutput(images=image)
|
||||
|
||||
|
||||
|
||||
|
||||
# TODO this is rough. Need to properly stack unconditional
|
||||
class FluxWithCFGPipeline(FluxPipeline):
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 7.0,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._joint_attention_kwargs = joint_attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
lora_scale = (
|
||||
self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None
|
||||
)
|
||||
(
|
||||
prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
text_ids,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
(
|
||||
negative_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
negative_text_ids,
|
||||
) = self.encode_prompt(
|
||||
prompt=negative_prompt,
|
||||
prompt_2=negative_prompt_2,
|
||||
prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
lora_scale=lora_scale,
|
||||
)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels // 4
|
||||
latents, latent_image_ids = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 5. Prepare timesteps
|
||||
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.base_image_seq_len,
|
||||
self.scheduler.config.max_image_seq_len,
|
||||
self.scheduler.config.base_shift,
|
||||
self.scheduler.config.max_shift,
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
# handle guidance
|
||||
if self.transformer.config.guidance_embeds:
|
||||
guidance = torch.tensor([guidance_scale], device=device)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
noise_pred_text = self.transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=pooled_prompt_embeds,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# todo combine these
|
||||
noise_pred_uncond = self.transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=negative_pooled_prompt_embeds,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
txt_ids=negative_text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
if torch.backends.mps.is_available():
|
||||
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
if output_type == "latent":
|
||||
image = latents
|
||||
|
||||
else:
|
||||
latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)
|
||||
latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = self.image_processor.postprocess(image, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return FluxPipelineOutput(images=image)
|
||||
@@ -19,10 +19,11 @@ class ACTION_TYPES_SLIDER:
|
||||
|
||||
|
||||
class PromptEmbeds:
|
||||
text_embeds: torch.Tensor
|
||||
pooled_embeds: Union[torch.Tensor, None]
|
||||
# text_embeds: torch.Tensor
|
||||
# pooled_embeds: Union[torch.Tensor, None]
|
||||
# attention_mask: Union[torch.Tensor, None]
|
||||
|
||||
def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Tensor]) -> None:
|
||||
def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Tensor], attention_mask=None) -> None:
|
||||
if isinstance(args, list) or isinstance(args, tuple):
|
||||
# xl
|
||||
self.text_embeds = args[0]
|
||||
@@ -32,23 +33,34 @@ class PromptEmbeds:
|
||||
self.text_embeds = args
|
||||
self.pooled_embeds = None
|
||||
|
||||
self.attention_mask = attention_mask
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.text_embeds = self.text_embeds.to(*args, **kwargs)
|
||||
if self.pooled_embeds is not None:
|
||||
self.pooled_embeds = self.pooled_embeds.to(*args, **kwargs)
|
||||
if self.attention_mask is not None:
|
||||
self.attention_mask = self.attention_mask.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def detach(self):
|
||||
self.text_embeds = self.text_embeds.detach()
|
||||
if self.pooled_embeds is not None:
|
||||
self.pooled_embeds = self.pooled_embeds.detach()
|
||||
return self
|
||||
new_embeds = self.clone()
|
||||
new_embeds.text_embeds = new_embeds.text_embeds.detach()
|
||||
if new_embeds.pooled_embeds is not None:
|
||||
new_embeds.pooled_embeds = new_embeds.pooled_embeds.detach()
|
||||
if new_embeds.attention_mask is not None:
|
||||
new_embeds.attention_mask = new_embeds.attention_mask.detach()
|
||||
return new_embeds
|
||||
|
||||
def clone(self):
|
||||
if self.pooled_embeds is not None:
|
||||
return PromptEmbeds([self.text_embeds.clone(), self.pooled_embeds.clone()])
|
||||
prompt_embeds = PromptEmbeds([self.text_embeds.clone(), self.pooled_embeds.clone()])
|
||||
else:
|
||||
return PromptEmbeds(self.text_embeds.clone())
|
||||
prompt_embeds = PromptEmbeds(self.text_embeds.clone())
|
||||
|
||||
if self.attention_mask is not None:
|
||||
prompt_embeds.attention_mask = self.attention_mask.clone()
|
||||
return prompt_embeds
|
||||
|
||||
|
||||
class EncodedPromptPair:
|
||||
|
||||
410
toolkit/reference_adapter.py
Normal file
410
toolkit/reference_adapter.py
Normal file
@@ -0,0 +1,410 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import sys
|
||||
|
||||
from PIL import Image
|
||||
from torch.nn import Parameter
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from toolkit.basic import adain
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from toolkit.saving import load_ip_adapter_model
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from typing import TYPE_CHECKING, Union, Iterator, Mapping, Any, Tuple, List, Optional, Dict
|
||||
from collections import OrderedDict
|
||||
from ipadapter.ip_adapter.attention_processor import AttnProcessor, IPAttnProcessor, IPAttnProcessor2_0, \
|
||||
AttnProcessor2_0
|
||||
from ipadapter.ip_adapter.ip_adapter import ImageProjModel
|
||||
from ipadapter.ip_adapter.resampler import Resampler
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
import weakref
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
from diffusers import (
|
||||
EulerDiscreteScheduler,
|
||||
DDPMScheduler,
|
||||
)
|
||||
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPVisionModelWithProjection
|
||||
)
|
||||
from toolkit.models.size_agnostic_feature_encoder import SAFEImageProcessor, SAFEVisionModel
|
||||
|
||||
from transformers import ViTHybridImageProcessor, ViTHybridForImageClassification
|
||||
|
||||
from transformers import ViTFeatureExtractor, ViTForImageClassification
|
||||
|
||||
import torch.nn.functional as F
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ReferenceAttnProcessor2_0(torch.nn.Module):
|
||||
r"""
|
||||
Attention processor for IP-Adapater for PyTorch 2.0.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
scale (`float`, defaults to 1.0):
|
||||
the weight scale of image prompt.
|
||||
num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16):
|
||||
The context length of the image features.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4, adapter=None):
|
||||
super().__init__()
|
||||
|
||||
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.hidden_size = hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
self.ref_net = nn.Linear(hidden_size, hidden_size)
|
||||
self.blend = nn.Parameter(torch.zeros(hidden_size))
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self._memory = None
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
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)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
if self.adapter_ref().is_active:
|
||||
if self.adapter_ref().reference_mode == "write":
|
||||
# write_mode
|
||||
memory_ref = self.ref_net(hidden_states)
|
||||
self._memory = memory_ref
|
||||
elif self.adapter_ref().reference_mode == "read":
|
||||
# read_mode
|
||||
if self._memory is None:
|
||||
print("Warning: no memory to read from")
|
||||
else:
|
||||
|
||||
saved_hidden_states = self._memory
|
||||
try:
|
||||
new_hidden_states = saved_hidden_states
|
||||
blend = self.blend
|
||||
# expand the blend buyt keep dim 0 the same (batch)
|
||||
while blend.ndim < new_hidden_states.ndim:
|
||||
blend = blend.unsqueeze(0)
|
||||
# expand batch
|
||||
blend = torch.cat([blend] * new_hidden_states.shape[0], dim=0)
|
||||
hidden_states = blend * new_hidden_states + (1 - blend) * hidden_states
|
||||
except Exception as e:
|
||||
raise Exception(f"Error blending: {e}")
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ReferenceAdapter(torch.nn.Module):
|
||||
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: 'AdapterConfig'):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.reference_mode = "read"
|
||||
self.current_scale = 1.0
|
||||
self.is_active = True
|
||||
self._reference_images = None
|
||||
self._reference_latents = None
|
||||
self.has_memory = False
|
||||
|
||||
self.noise_scheduler: Union[DDPMScheduler, EulerDiscreteScheduler] = None
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
unet_sd = sd.unet.state_dict()
|
||||
for name in sd.unet.attn_processors.keys():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else sd.unet.config['cross_attention_dim']
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = sd.unet.config['block_out_channels'][-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(sd.unet.config['block_out_channels']))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
else:
|
||||
# they didnt have this, but would lead to undefined below
|
||||
raise ValueError(f"unknown attn processor name: {name}")
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AttnProcessor2_0()
|
||||
else:
|
||||
# layer_name = name.split(".processor")[0]
|
||||
# weights = {
|
||||
# "to_k_ip.weight": unet_sd[layer_name + ".to_k.weight"],
|
||||
# "to_v_ip.weight": unet_sd[layer_name + ".to_v.weight"],
|
||||
# }
|
||||
|
||||
attn_procs[name] = ReferenceAttnProcessor2_0(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
num_tokens=self.config.num_tokens,
|
||||
adapter=self
|
||||
)
|
||||
# attn_procs[name].load_state_dict(weights)
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
|
||||
sd.adapter = self
|
||||
self.unet_ref: weakref.ref = weakref.ref(sd.unet)
|
||||
self.adapter_modules = adapter_modules
|
||||
# load the weights if we have some
|
||||
if self.config.name_or_path:
|
||||
loaded_state_dict = load_ip_adapter_model(
|
||||
self.config.name_or_path,
|
||||
device='cpu',
|
||||
dtype=sd.torch_dtype
|
||||
)
|
||||
self.load_state_dict(loaded_state_dict)
|
||||
|
||||
self.set_scale(1.0)
|
||||
self.attach()
|
||||
self.to(self.device, self.sd_ref().torch_dtype)
|
||||
|
||||
# if self.config.train_image_encoder:
|
||||
# self.image_encoder.train()
|
||||
# self.image_encoder.requires_grad_(True)
|
||||
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
super().to(*args, **kwargs)
|
||||
# self.image_encoder.to(*args, **kwargs)
|
||||
# self.image_proj_model.to(*args, **kwargs)
|
||||
self.adapter_modules.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
def load_reference_adapter(self, state_dict: Union[OrderedDict, dict]):
|
||||
reference_layers = torch.nn.ModuleList(self.pipe.unet.attn_processors.values())
|
||||
reference_layers.load_state_dict(state_dict["reference_adapter"])
|
||||
|
||||
# def load_state_dict(self, state_dict: Union[OrderedDict, dict]):
|
||||
# self.load_ip_adapter(state_dict)
|
||||
|
||||
def state_dict(self) -> OrderedDict:
|
||||
state_dict = OrderedDict()
|
||||
state_dict["reference_adapter"] = self.adapter_modules.state_dict()
|
||||
return state_dict
|
||||
|
||||
def get_scale(self):
|
||||
return self.current_scale
|
||||
|
||||
def set_reference_images(self, reference_images: Optional[torch.Tensor]):
|
||||
self._reference_images = reference_images.clone().detach()
|
||||
self._reference_latents = None
|
||||
self.clear_memory()
|
||||
|
||||
def set_blank_reference_images(self, batch_size):
|
||||
self._reference_images = torch.zeros((batch_size, 3, 512, 512), device=self.device, dtype=self.sd_ref().torch_dtype)
|
||||
self._reference_latents = torch.zeros((batch_size, 4, 64, 64), device=self.device, dtype=self.sd_ref().torch_dtype)
|
||||
self.clear_memory()
|
||||
|
||||
|
||||
def set_scale(self, scale):
|
||||
self.current_scale = scale
|
||||
for attn_processor in self.sd_ref().unet.attn_processors.values():
|
||||
if isinstance(attn_processor, ReferenceAttnProcessor2_0):
|
||||
attn_processor.scale = scale
|
||||
|
||||
|
||||
def attach(self):
|
||||
unet = self.sd_ref().unet
|
||||
self._original_unet_forward = unet.forward
|
||||
unet.forward = lambda *args, **kwargs: self.unet_forward(*args, **kwargs)
|
||||
if self.sd_ref().network is not None:
|
||||
# set network to not merge in
|
||||
self.sd_ref().network.can_merge_in = False
|
||||
|
||||
def unet_forward(self, sample, timestep, encoder_hidden_states, *args, **kwargs):
|
||||
skip = False
|
||||
if self._reference_images is None and self._reference_latents is None:
|
||||
skip = True
|
||||
if not self.is_active:
|
||||
skip = True
|
||||
|
||||
if self.has_memory:
|
||||
skip = True
|
||||
|
||||
if not skip:
|
||||
if self.sd_ref().network is not None:
|
||||
self.sd_ref().network.is_active = True
|
||||
if self.sd_ref().network.is_merged_in:
|
||||
raise ValueError("network is merged in, but we are not supposed to be merged in")
|
||||
# send it through our forward first
|
||||
self.forward(sample, timestep, encoder_hidden_states, *args, **kwargs)
|
||||
|
||||
if self.sd_ref().network is not None:
|
||||
self.sd_ref().network.is_active = False
|
||||
|
||||
# Send it through the original unet forward
|
||||
return self._original_unet_forward(sample, timestep, encoder_hidden_states, args, **kwargs)
|
||||
|
||||
|
||||
# use drop for prompt dropout, or negatives
|
||||
def forward(self, sample, timestep, encoder_hidden_states, *args, **kwargs):
|
||||
if not self.noise_scheduler:
|
||||
raise ValueError("noise scheduler not set")
|
||||
if not self.is_active or (self._reference_images is None and self._reference_latents is None):
|
||||
raise ValueError("reference adapter not active or no reference images set")
|
||||
# todo may need to handle cfg?
|
||||
self.reference_mode = "write"
|
||||
|
||||
if self._reference_latents is None:
|
||||
self._reference_latents = self.sd_ref().encode_images(self._reference_images.to(
|
||||
self.device, self.sd_ref().torch_dtype
|
||||
)).detach()
|
||||
# create a sample from our reference images
|
||||
reference_latents = self._reference_latents.clone().detach().to(self.device, self.sd_ref().torch_dtype)
|
||||
# if our num of samples are half of incoming, we are doing cfg. Zero out the first half (unconditional)
|
||||
if reference_latents.shape[0] * 2 == sample.shape[0]:
|
||||
# we are doing cfg
|
||||
# Unconditional goes first
|
||||
reference_latents = torch.cat([torch.zeros_like(reference_latents), reference_latents], dim=0).detach()
|
||||
|
||||
# resize it so reference_latents will fit inside sample in the center
|
||||
width_scale = sample.shape[2] / reference_latents.shape[2]
|
||||
height_scale = sample.shape[3] / reference_latents.shape[3]
|
||||
scale = min(width_scale, height_scale)
|
||||
# resize the reference latents
|
||||
|
||||
mode = "bilinear" if scale > 1.0 else "bicubic"
|
||||
|
||||
reference_latents = F.interpolate(
|
||||
reference_latents,
|
||||
size=(int(reference_latents.shape[2] * scale), int(reference_latents.shape[3] * scale)),
|
||||
mode=mode,
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# add 0 padding if needed
|
||||
width_pad = (sample.shape[2] - reference_latents.shape[2]) / 2
|
||||
height_pad = (sample.shape[3] - reference_latents.shape[3]) / 2
|
||||
reference_latents = F.pad(
|
||||
reference_latents,
|
||||
(math.floor(width_pad), math.floor(width_pad), math.ceil(height_pad), math.ceil(height_pad)),
|
||||
mode="constant",
|
||||
value=0
|
||||
)
|
||||
|
||||
# resize again just to make sure it is exact same size
|
||||
reference_latents = F.interpolate(
|
||||
reference_latents,
|
||||
size=(sample.shape[2], sample.shape[3]),
|
||||
mode="bicubic",
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# todo maybe add same noise to the sample? For now we will send it through with no noise
|
||||
# sample_imgs = self.noise_scheduler.add_noise(sample_imgs, timestep)
|
||||
self._original_unet_forward(reference_latents, timestep, encoder_hidden_states, *args, **kwargs)
|
||||
self.reference_mode = "read"
|
||||
self.has_memory = True
|
||||
return None
|
||||
|
||||
def parameters(self, recurse: bool = True) -> Iterator[Parameter]:
|
||||
for attn_processor in self.adapter_modules:
|
||||
yield from attn_processor.parameters(recurse)
|
||||
# yield from self.image_proj_model.parameters(recurse)
|
||||
# if self.config.train_image_encoder:
|
||||
# yield from self.image_encoder.parameters(recurse)
|
||||
# if self.config.train_image_encoder:
|
||||
# yield from self.image_encoder.parameters(recurse)
|
||||
# self.image_encoder.train()
|
||||
# else:
|
||||
# for attn_processor in self.adapter_modules:
|
||||
# yield from attn_processor.parameters(recurse)
|
||||
# yield from self.image_proj_model.parameters(recurse)
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
|
||||
strict = False
|
||||
# self.image_proj_model.load_state_dict(state_dict["image_proj"], strict=strict)
|
||||
self.adapter_modules.load_state_dict(state_dict["reference_adapter"], strict=strict)
|
||||
|
||||
def enable_gradient_checkpointing(self):
|
||||
self.image_encoder.gradient_checkpointing = True
|
||||
|
||||
def clear_memory(self):
|
||||
for attn_processor in self.adapter_modules:
|
||||
if isinstance(attn_processor, ReferenceAttnProcessor2_0):
|
||||
attn_processor._memory = None
|
||||
self.has_memory = False
|
||||
160
toolkit/resampler.py
Normal file
160
toolkit/resampler.py
Normal file
@@ -0,0 +1,160 @@
|
||||
# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py
|
||||
# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py
|
||||
# and https://github.com/tencent-ailab/IP-Adapter/blob/9fc189e3fb389cc2b60a7d0c0850e083a716ea6e/ip_adapter/resampler.py
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head ** -0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=8,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_queries=8,
|
||||
embedding_dim=768,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
max_seq_len: int = 257, # CLIP tokens + CLS token
|
||||
apply_pos_emb: bool = False,
|
||||
num_latents_mean_pooled: int = 0,
|
||||
# number of latents derived from mean pooled representation of the sequence
|
||||
):
|
||||
super().__init__()
|
||||
self.pos_emb = nn.Embedding(max_seq_len, embedding_dim) if apply_pos_emb else None
|
||||
|
||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim ** 0.5)
|
||||
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.to_latents_from_mean_pooled_seq = (
|
||||
nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, dim * num_latents_mean_pooled),
|
||||
Rearrange("b (n d) -> b n d", n=num_latents_mean_pooled),
|
||||
)
|
||||
if num_latents_mean_pooled > 0
|
||||
else None
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
if self.pos_emb is not None:
|
||||
n, device = x.shape[1], x.device
|
||||
pos_emb = self.pos_emb(torch.arange(n, device=device))
|
||||
x = x + pos_emb
|
||||
|
||||
latents = self.latents.repeat(x.size(0), 1, 1)
|
||||
|
||||
x = self.proj_in(x)
|
||||
|
||||
if self.to_latents_from_mean_pooled_seq:
|
||||
meanpooled_seq = masked_mean(x, dim=1, mask=torch.ones(x.shape[:2], device=x.device, dtype=torch.bool))
|
||||
meanpooled_latents = self.to_latents_from_mean_pooled_seq(meanpooled_seq)
|
||||
latents = torch.cat((meanpooled_latents, latents), dim=-2)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
return self.norm_out(latents)
|
||||
|
||||
|
||||
def masked_mean(t, *, dim, mask=None):
|
||||
if mask is None:
|
||||
return t.mean(dim=dim)
|
||||
|
||||
denom = mask.sum(dim=dim, keepdim=True)
|
||||
mask = rearrange(mask, "b n -> b n 1")
|
||||
masked_t = t.masked_fill(~mask, 0.0)
|
||||
|
||||
return masked_t.sum(dim=dim) / denom.clamp(min=1e-5)
|
||||
@@ -1,4 +1,5 @@
|
||||
import copy
|
||||
import math
|
||||
|
||||
from diffusers import (
|
||||
DDPMScheduler,
|
||||
@@ -12,9 +13,12 @@ from diffusers import (
|
||||
HeunDiscreteScheduler,
|
||||
KDPM2DiscreteScheduler,
|
||||
KDPM2AncestralDiscreteScheduler,
|
||||
LCMScheduler
|
||||
LCMScheduler,
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
|
||||
from k_diffusion.external import CompVisDenoiser
|
||||
|
||||
from toolkit.samplers.custom_lcm_scheduler import CustomLCMScheduler
|
||||
@@ -25,9 +29,9 @@ SCHEDULER_LINEAR_END = 0.0120
|
||||
SCHEDULER_TIMESTEPS = 1000
|
||||
SCHEDLER_SCHEDULE = "scaled_linear"
|
||||
|
||||
sdxl_sampler_config = {
|
||||
"_class_name": "EulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.19.0.dev0",
|
||||
sd_config = {
|
||||
"_class_name": "EulerAncestralDiscreteScheduler",
|
||||
"_diffusers_version": "0.24.0.dev0",
|
||||
"beta_end": 0.012,
|
||||
"beta_schedule": "scaled_linear",
|
||||
"beta_start": 0.00085,
|
||||
@@ -37,18 +41,52 @@ sdxl_sampler_config = {
|
||||
"prediction_type": "epsilon",
|
||||
"sample_max_value": 1.0,
|
||||
"set_alpha_to_one": False,
|
||||
# "skip_prk_steps": False, # for training
|
||||
"skip_prk_steps": True,
|
||||
"steps_offset": 1,
|
||||
# "steps_offset": 1,
|
||||
"steps_offset": 0,
|
||||
# "timestep_spacing": "trailing", # for training
|
||||
"timestep_spacing": "leading",
|
||||
"trained_betas": None,
|
||||
"use_karras_sigmas": False
|
||||
"trained_betas": None
|
||||
}
|
||||
|
||||
pixart_config = {
|
||||
"_class_name": "DPMSolverMultistepScheduler",
|
||||
"_diffusers_version": "0.22.0.dev0",
|
||||
"algorithm_type": "dpmsolver++",
|
||||
"beta_end": 0.02,
|
||||
"beta_schedule": "linear",
|
||||
"beta_start": 0.0001,
|
||||
"dynamic_thresholding_ratio": 0.995,
|
||||
"euler_at_final": False,
|
||||
# "lambda_min_clipped": -Infinity,
|
||||
"lambda_min_clipped": -math.inf,
|
||||
"lower_order_final": True,
|
||||
"num_train_timesteps": 1000,
|
||||
"prediction_type": "epsilon",
|
||||
"sample_max_value": 1.0,
|
||||
"solver_order": 2,
|
||||
"solver_type": "midpoint",
|
||||
"steps_offset": 0,
|
||||
"thresholding": False,
|
||||
"timestep_spacing": "linspace",
|
||||
"trained_betas": None,
|
||||
"use_karras_sigmas": False,
|
||||
"use_lu_lambdas": False,
|
||||
"variance_type": None
|
||||
}
|
||||
|
||||
|
||||
def get_sampler(
|
||||
sampler: str,
|
||||
kwargs: dict = None,
|
||||
arch: str = "sd"
|
||||
):
|
||||
sched_init_args = {}
|
||||
if kwargs is not None:
|
||||
sched_init_args.update(kwargs)
|
||||
|
||||
config_to_use = copy.deepcopy(sd_config) if arch == "sd" else copy.deepcopy(pixart_config)
|
||||
|
||||
if sampler.startswith("k_"):
|
||||
sched_init_args["use_karras_sigmas"] = True
|
||||
@@ -80,13 +118,23 @@ def get_sampler(
|
||||
scheduler_cls = LCMScheduler
|
||||
elif sampler == "custom_lcm":
|
||||
scheduler_cls = CustomLCMScheduler
|
||||
elif sampler == "flowmatch":
|
||||
scheduler_cls = CustomFlowMatchEulerDiscreteScheduler
|
||||
config_to_use = {
|
||||
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.29.0.dev0",
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Sampler {sampler} not supported")
|
||||
|
||||
config = copy.deepcopy(sdxl_sampler_config)
|
||||
|
||||
config = copy.deepcopy(config_to_use)
|
||||
config.update(sched_init_args)
|
||||
|
||||
scheduler = scheduler_cls.from_config(config)
|
||||
|
||||
|
||||
return scheduler
|
||||
|
||||
|
||||
|
||||
100
toolkit/samplers/custom_flowmatch_sampler.py
Normal file
100
toolkit/samplers/custom_flowmatch_sampler.py
Normal file
@@ -0,0 +1,100 @@
|
||||
import math
|
||||
from typing import Union
|
||||
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
import torch
|
||||
|
||||
|
||||
class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.init_noise_sigma = 1.0
|
||||
|
||||
with torch.no_grad():
|
||||
# create weights for timesteps
|
||||
num_timesteps = 1000
|
||||
# Bell-Shaped Mean-Normalized Timestep Weighting
|
||||
# bsmntw? need a better name
|
||||
|
||||
x = torch.arange(num_timesteps, dtype=torch.float32)
|
||||
y = torch.exp(-2 * ((x - num_timesteps / 2) / num_timesteps) ** 2)
|
||||
|
||||
# Shift minimum to 0
|
||||
y_shifted = y - y.min()
|
||||
|
||||
# Scale to make mean 1
|
||||
bsmntw_weighing = y_shifted * (num_timesteps / y_shifted.sum())
|
||||
|
||||
# Create linear timesteps from 1000 to 0
|
||||
timesteps = torch.linspace(1000, 0, num_timesteps, device='cpu')
|
||||
|
||||
self.linear_timesteps = timesteps
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
pass
|
||||
|
||||
def get_weights_for_timesteps(self, timesteps: torch.Tensor) -> torch.Tensor:
|
||||
# Get the indices of the timesteps
|
||||
step_indices = [(self.timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
# Get the weights for the timesteps
|
||||
weights = self.linear_timesteps_weights[step_indices].flatten()
|
||||
|
||||
return weights
|
||||
|
||||
def get_sigmas(self, timesteps: torch.Tensor, n_dim, dtype, device) -> torch.Tensor:
|
||||
sigmas = self.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = self.timesteps.to(device)
|
||||
timesteps = timesteps.to(device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
return sigma
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
## ref https://github.com/huggingface/diffusers/blob/fbe29c62984c33c6cf9cf7ad120a992fe6d20854/examples/dreambooth/train_dreambooth_sd3.py#L1578
|
||||
## Add noise according to flow matching.
|
||||
## zt = (1 - texp) * x + texp * z1
|
||||
|
||||
# sigmas = get_sigmas(timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
# noisy_model_input = (1.0 - sigmas) * model_input + sigmas * noise
|
||||
|
||||
# timestep needs to be in [0, 1], we store them in [0, 1000]
|
||||
# noisy_sample = (1 - timestep) * latent + timestep * noise
|
||||
t_01 = (timesteps / 1000).to(original_samples.device)
|
||||
noisy_model_input = (1 - t_01) * original_samples + t_01 * noise
|
||||
|
||||
# n_dim = original_samples.ndim
|
||||
# sigmas = self.get_sigmas(timesteps, n_dim, original_samples.dtype, original_samples.device)
|
||||
# noisy_model_input = (1.0 - sigmas) * original_samples + sigmas * noise
|
||||
return noisy_model_input
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, timestep: Union[float, torch.Tensor]) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def set_train_timesteps(self, num_timesteps, device, linear=False):
|
||||
if linear:
|
||||
timesteps = torch.linspace(1000, 0, num_timesteps, device=device)
|
||||
self.timesteps = timesteps
|
||||
return timesteps
|
||||
else:
|
||||
# distribute them closer to center. Inference distributes them as a bias toward first
|
||||
# Generate values from 0 to 1
|
||||
t = torch.sigmoid(torch.randn((num_timesteps,), device=device))
|
||||
|
||||
# Scale and reverse the values to go from 1000 to 0
|
||||
timesteps = ((1 - t) * 1000)
|
||||
|
||||
# Sort the timesteps in descending order
|
||||
timesteps, _ = torch.sort(timesteps, descending=True)
|
||||
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
|
||||
return timesteps
|
||||
@@ -97,7 +97,7 @@ def convert_state_dict_to_ldm_with_mapping(
|
||||
|
||||
def get_ldm_state_dict_from_diffusers(
|
||||
state_dict: 'OrderedDict',
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'sdxl_refiner'] = '2',
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega', 'sdxl_refiner'] = '2',
|
||||
device='cpu',
|
||||
dtype=get_torch_dtype('fp32'),
|
||||
):
|
||||
@@ -115,6 +115,10 @@ def get_ldm_state_dict_from_diffusers(
|
||||
# load our base
|
||||
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_ssd_ldm_base.safetensors')
|
||||
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_ssd.json')
|
||||
elif sd_version == 'vega':
|
||||
# load our base
|
||||
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_vega_ldm_base.safetensors')
|
||||
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_vega.json')
|
||||
elif sd_version == 'sdxl_refiner':
|
||||
# load our base
|
||||
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_refiner_ldm_base.safetensors')
|
||||
@@ -137,7 +141,7 @@ def save_ldm_model_from_diffusers(
|
||||
output_file: str,
|
||||
meta: 'OrderedDict',
|
||||
save_dtype=get_torch_dtype('fp16'),
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd'] = '2'
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega'] = '2'
|
||||
):
|
||||
converted_state_dict = get_ldm_state_dict_from_diffusers(
|
||||
sd.state_dict(),
|
||||
@@ -156,11 +160,11 @@ def save_lora_from_diffusers(
|
||||
output_file: str,
|
||||
meta: 'OrderedDict',
|
||||
save_dtype=get_torch_dtype('fp16'),
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd'] = '2'
|
||||
sd_version: Literal['1', '2', 'sdxl', 'ssd', 'vega'] = '2'
|
||||
):
|
||||
converted_state_dict = OrderedDict()
|
||||
# only handle sxdxl for now
|
||||
if sd_version != 'sdxl' and sd_version != 'ssd':
|
||||
if sd_version != 'sdxl' and sd_version != 'ssd' and sd_version != 'vega':
|
||||
raise ValueError(f"Invalid sd_version {sd_version}")
|
||||
for key, value in lora_state_dict.items():
|
||||
# todo verify if this works with ssd
|
||||
@@ -204,19 +208,24 @@ def load_t2i_model(
|
||||
return converted_state_dict
|
||||
|
||||
|
||||
IP_ADAPTER_MODULES = ['image_proj', 'ip_adapter']
|
||||
|
||||
|
||||
def save_ip_adapter_from_diffusers(
|
||||
combined_state_dict: 'OrderedDict',
|
||||
output_file: str,
|
||||
meta: 'OrderedDict',
|
||||
dtype=get_torch_dtype('fp16'),
|
||||
direct_save: bool = False
|
||||
):
|
||||
# todo: test compatibility with non diffusers
|
||||
|
||||
converted_state_dict = OrderedDict()
|
||||
for module_name, state_dict in combined_state_dict.items():
|
||||
for key, value in state_dict.items():
|
||||
converted_state_dict[f"{module_name}.{key}"] = value.detach().to('cpu', dtype=dtype)
|
||||
if direct_save:
|
||||
converted_state_dict[module_name] = state_dict.detach().to('cpu', dtype=dtype)
|
||||
else:
|
||||
for key, value in state_dict.items():
|
||||
converted_state_dict[f"{module_name}.{key}"] = value.detach().to('cpu', dtype=dtype)
|
||||
|
||||
# make sure parent folder exists
|
||||
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
||||
@@ -226,12 +235,15 @@ def save_ip_adapter_from_diffusers(
|
||||
def load_ip_adapter_model(
|
||||
path_to_file,
|
||||
device: Union[str] = 'cpu',
|
||||
dtype: torch.dtype = torch.float32
|
||||
dtype: torch.dtype = torch.float32,
|
||||
direct_load: bool = False
|
||||
):
|
||||
# check if it is safetensors or checkpoint
|
||||
if path_to_file.endswith('.safetensors'):
|
||||
raw_state_dict = load_file(path_to_file, device)
|
||||
combined_state_dict = OrderedDict()
|
||||
if direct_load:
|
||||
return raw_state_dict
|
||||
for combo_key, value in raw_state_dict.items():
|
||||
key_split = combo_key.split('.')
|
||||
module_name = key_split.pop(0)
|
||||
@@ -241,3 +253,78 @@ def load_ip_adapter_model(
|
||||
return combined_state_dict
|
||||
else:
|
||||
return torch.load(path_to_file, map_location=device)
|
||||
|
||||
def load_custom_adapter_model(
|
||||
path_to_file,
|
||||
device: Union[str] = 'cpu',
|
||||
dtype: torch.dtype = torch.float32
|
||||
):
|
||||
# check if it is safetensors or checkpoint
|
||||
if path_to_file.endswith('.safetensors'):
|
||||
raw_state_dict = load_file(path_to_file, device)
|
||||
combined_state_dict = OrderedDict()
|
||||
device = device if isinstance(device, torch.device) else torch.device(device)
|
||||
dtype = dtype if isinstance(dtype, torch.dtype) else get_torch_dtype(dtype)
|
||||
for combo_key, value in raw_state_dict.items():
|
||||
key_split = combo_key.split('.')
|
||||
module_name = key_split.pop(0)
|
||||
if module_name not in combined_state_dict:
|
||||
combined_state_dict[module_name] = OrderedDict()
|
||||
combined_state_dict[module_name]['.'.join(key_split)] = value.detach().to(device, dtype=dtype)
|
||||
return combined_state_dict
|
||||
else:
|
||||
return torch.load(path_to_file, map_location=device)
|
||||
|
||||
|
||||
def get_lora_keymap_from_model_keymap(model_keymap: 'OrderedDict') -> 'OrderedDict':
|
||||
lora_keymap = OrderedDict()
|
||||
|
||||
# see if we have dual text encoders " a key that starts with conditioner.embedders.1
|
||||
has_dual_text_encoders = False
|
||||
for key in model_keymap:
|
||||
if key.startswith('conditioner.embedders.1'):
|
||||
has_dual_text_encoders = True
|
||||
break
|
||||
# map through the keys and values
|
||||
for key, value in model_keymap.items():
|
||||
# ignore bias weights
|
||||
if key.endswith('bias'):
|
||||
continue
|
||||
if key.endswith('.weight'):
|
||||
# remove the .weight
|
||||
key = key[:-7]
|
||||
if value.endswith(".weight"):
|
||||
# remove the .weight
|
||||
value = value[:-7]
|
||||
|
||||
# unet for all
|
||||
key = key.replace('model.diffusion_model', 'lora_unet')
|
||||
if value.startswith('unet'):
|
||||
value = f"lora_{value}"
|
||||
|
||||
# text encoder
|
||||
if has_dual_text_encoders:
|
||||
key = key.replace('conditioner.embedders.0', 'lora_te1')
|
||||
key = key.replace('conditioner.embedders.1', 'lora_te2')
|
||||
if value.startswith('te0') or value.startswith('te1'):
|
||||
value = f"lora_{value}"
|
||||
value.replace('lora_te1', 'lora_te2')
|
||||
value.replace('lora_te0', 'lora_te1')
|
||||
|
||||
key = key.replace('cond_stage_model.transformer', 'lora_te')
|
||||
|
||||
if value.startswith('te_'):
|
||||
value = f"lora_{value}"
|
||||
|
||||
# replace periods with underscores
|
||||
key = key.replace('.', '_')
|
||||
value = value.replace('.', '_')
|
||||
|
||||
# add all the weights
|
||||
lora_keymap[f"{key}.lora_down.weight"] = f"{value}.lora_down.weight"
|
||||
lora_keymap[f"{key}.lora_down.bias"] = f"{value}.lora_down.bias"
|
||||
lora_keymap[f"{key}.lora_up.weight"] = f"{value}.lora_up.weight"
|
||||
lora_keymap[f"{key}.lora_up.bias"] = f"{value}.lora_up.bias"
|
||||
lora_keymap[f"{key}.alpha"] = f"{value}.alpha"
|
||||
|
||||
return lora_keymap
|
||||
|
||||
@@ -26,7 +26,7 @@ def get_lr_scheduler(
|
||||
optimizer, **kwargs
|
||||
)
|
||||
elif name == "constant":
|
||||
if 'facor' not in kwargs:
|
||||
if 'factor' not in kwargs:
|
||||
kwargs['factor'] = 1.0
|
||||
|
||||
return torch.optim.lr_scheduler.ConstantLR(optimizer, **kwargs)
|
||||
|
||||
@@ -84,5 +84,8 @@ def get_train_sd_device_state_preset(
|
||||
preset['adapter']['training'] = True
|
||||
preset['adapter']['device'] = device
|
||||
preset['unet']['training'] = True
|
||||
preset['unet']['requires_grad'] = False
|
||||
preset['unet']['device'] = device
|
||||
preset['text_encoder']['device'] = device
|
||||
|
||||
return preset
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -158,7 +158,7 @@ def get_style_model_and_losses(
|
||||
):
|
||||
# content_layers = ['conv_4']
|
||||
# style_layers = ['conv_1', 'conv_2', 'conv_3', 'conv_4', 'conv_5']
|
||||
content_layers = ['conv2_2', 'conv3_2', 'conv4_2', 'conv5_2']
|
||||
content_layers = ['conv2_2', 'conv3_2', 'conv4_2']
|
||||
style_layers = ['conv2_1', 'conv3_1', 'conv4_1']
|
||||
cnn = models.vgg19(pretrained=True).features.to(device, dtype=dtype).eval()
|
||||
# set all weights in the model to our dtype
|
||||
|
||||
@@ -63,3 +63,29 @@ class Timer:
|
||||
else:
|
||||
# There was an exception, cancel the timer
|
||||
self.cancel(self.current_timer)
|
||||
|
||||
|
||||
class DummyTimer:
|
||||
def __init__(self, name='Timer'):
|
||||
self.name = name
|
||||
|
||||
def start(self, timer_name):
|
||||
pass
|
||||
|
||||
def stop(self, timer_name):
|
||||
pass
|
||||
|
||||
def print(self):
|
||||
pass
|
||||
|
||||
def reset(self):
|
||||
pass
|
||||
|
||||
def __call__(self, timer_name):
|
||||
return self
|
||||
|
||||
def __enter__(self):
|
||||
pass
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
pass
|
||||
|
||||
@@ -3,7 +3,7 @@ import hashlib
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Union
|
||||
from typing import TYPE_CHECKING, Union, List
|
||||
import sys
|
||||
|
||||
from torch.cuda.amp import GradScaler
|
||||
@@ -13,7 +13,6 @@ from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from diffusers import (
|
||||
StableDiffusionPipeline,
|
||||
DDPMScheduler,
|
||||
EulerAncestralDiscreteScheduler,
|
||||
DPMSolverMultistepScheduler,
|
||||
@@ -24,11 +23,11 @@ from diffusers import (
|
||||
EulerDiscreteScheduler,
|
||||
HeunDiscreteScheduler,
|
||||
KDPM2DiscreteScheduler,
|
||||
KDPM2AncestralDiscreteScheduler,
|
||||
KDPM2AncestralDiscreteScheduler
|
||||
)
|
||||
from library.lpw_stable_diffusion import StableDiffusionLongPromptWeightingPipeline
|
||||
import torch
|
||||
import re
|
||||
from transformers import T5Tokenizer, T5EncoderModel, UMT5EncoderModel
|
||||
|
||||
SCHEDULER_LINEAR_START = 0.00085
|
||||
SCHEDULER_LINEAR_END = 0.0120
|
||||
@@ -50,6 +49,8 @@ def get_torch_dtype(dtype_str):
|
||||
return torch.float16
|
||||
if dtype_str == "bf16" or dtype_str == "bfloat16":
|
||||
return torch.bfloat16
|
||||
if dtype_str == "8bit" or dtype_str == "e4m3fn" or dtype_str == "float8":
|
||||
return torch.float8_e4m3fn
|
||||
return dtype_str
|
||||
|
||||
|
||||
@@ -132,264 +133,9 @@ def match_noise_to_target_mean_offset(noise, target, mix=0.5, dim=None):
|
||||
return noise
|
||||
|
||||
|
||||
def sample_images(
|
||||
accelerator,
|
||||
args: argparse.Namespace,
|
||||
epoch,
|
||||
steps,
|
||||
device,
|
||||
vae,
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
unet,
|
||||
prompt_replacement=None,
|
||||
force_sample=False
|
||||
):
|
||||
"""
|
||||
StableDiffusionLongPromptWeightingPipelineの改造版を使うようにしたので、clip skipおよびプロンプトの重みづけに対応した
|
||||
"""
|
||||
if not force_sample:
|
||||
if args.sample_every_n_steps is None and args.sample_every_n_epochs is None:
|
||||
return
|
||||
if args.sample_every_n_epochs is not None:
|
||||
# sample_every_n_steps は無視する
|
||||
if epoch is None or epoch % args.sample_every_n_epochs != 0:
|
||||
return
|
||||
else:
|
||||
if steps % args.sample_every_n_steps != 0 or epoch is not None: # steps is not divisible or end of epoch
|
||||
return
|
||||
|
||||
is_sample_only = args.sample_only
|
||||
is_generating_only = hasattr(args, "is_generating_only") and args.is_generating_only
|
||||
|
||||
print(f"\ngenerating sample images at step / サンプル画像生成 ステップ: {steps}")
|
||||
if not os.path.isfile(args.sample_prompts):
|
||||
print(f"No prompt file / プロンプトファイルがありません: {args.sample_prompts}")
|
||||
return
|
||||
|
||||
org_vae_device = vae.device # CPUにいるはず
|
||||
vae.to(device)
|
||||
|
||||
# read prompts
|
||||
|
||||
# with open(args.sample_prompts, "rt", encoding="utf-8") as f:
|
||||
# prompts = f.readlines()
|
||||
|
||||
if args.sample_prompts.endswith(".txt"):
|
||||
with open(args.sample_prompts, "r", encoding="utf-8") as f:
|
||||
lines = f.readlines()
|
||||
prompts = [line.strip() for line in lines if len(line.strip()) > 0 and line[0] != "#"]
|
||||
elif args.sample_prompts.endswith(".json"):
|
||||
with open(args.sample_prompts, "r", encoding="utf-8") as f:
|
||||
prompts = json.load(f)
|
||||
|
||||
# schedulerを用意する
|
||||
sched_init_args = {}
|
||||
if args.sample_sampler == "ddim":
|
||||
scheduler_cls = DDIMScheduler
|
||||
elif args.sample_sampler == "ddpm": # ddpmはおかしくなるのでoptionから外してある
|
||||
scheduler_cls = DDPMScheduler
|
||||
elif args.sample_sampler == "pndm":
|
||||
scheduler_cls = PNDMScheduler
|
||||
elif args.sample_sampler == "lms" or args.sample_sampler == "k_lms":
|
||||
scheduler_cls = LMSDiscreteScheduler
|
||||
elif args.sample_sampler == "euler" or args.sample_sampler == "k_euler":
|
||||
scheduler_cls = EulerDiscreteScheduler
|
||||
elif args.sample_sampler == "euler_a" or args.sample_sampler == "k_euler_a":
|
||||
scheduler_cls = EulerAncestralDiscreteScheduler
|
||||
elif args.sample_sampler == "dpmsolver" or args.sample_sampler == "dpmsolver++":
|
||||
scheduler_cls = DPMSolverMultistepScheduler
|
||||
sched_init_args["algorithm_type"] = args.sample_sampler
|
||||
elif args.sample_sampler == "dpmsingle":
|
||||
scheduler_cls = DPMSolverSinglestepScheduler
|
||||
elif args.sample_sampler == "heun":
|
||||
scheduler_cls = HeunDiscreteScheduler
|
||||
elif args.sample_sampler == "dpm_2" or args.sample_sampler == "k_dpm_2":
|
||||
scheduler_cls = KDPM2DiscreteScheduler
|
||||
elif args.sample_sampler == "dpm_2_a" or args.sample_sampler == "k_dpm_2_a":
|
||||
scheduler_cls = KDPM2AncestralDiscreteScheduler
|
||||
else:
|
||||
scheduler_cls = DDIMScheduler
|
||||
|
||||
if args.v_parameterization:
|
||||
sched_init_args["prediction_type"] = "v_prediction"
|
||||
|
||||
scheduler = scheduler_cls(
|
||||
num_train_timesteps=SCHEDULER_TIMESTEPS,
|
||||
beta_start=SCHEDULER_LINEAR_START,
|
||||
beta_end=SCHEDULER_LINEAR_END,
|
||||
beta_schedule=SCHEDLER_SCHEDULE,
|
||||
**sched_init_args,
|
||||
)
|
||||
|
||||
# clip_sample=Trueにする
|
||||
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is False:
|
||||
# print("set clip_sample to True")
|
||||
scheduler.config.clip_sample = True
|
||||
|
||||
pipeline = StableDiffusionLongPromptWeightingPipeline(
|
||||
text_encoder=text_encoder,
|
||||
vae=vae,
|
||||
unet=unet,
|
||||
tokenizer=tokenizer,
|
||||
scheduler=scheduler,
|
||||
clip_skip=args.clip_skip,
|
||||
safety_checker=None,
|
||||
feature_extractor=None,
|
||||
requires_safety_checker=False,
|
||||
)
|
||||
pipeline.to(device)
|
||||
|
||||
if is_generating_only:
|
||||
save_dir = args.output_dir
|
||||
else:
|
||||
save_dir = args.output_dir + "/sample"
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
|
||||
|
||||
with torch.no_grad():
|
||||
with accelerator.autocast():
|
||||
for i, prompt in enumerate(prompts):
|
||||
if not accelerator.is_main_process:
|
||||
continue
|
||||
|
||||
if isinstance(prompt, dict):
|
||||
negative_prompt = prompt.get("negative_prompt")
|
||||
sample_steps = prompt.get("sample_steps", 30)
|
||||
width = prompt.get("width", 512)
|
||||
height = prompt.get("height", 512)
|
||||
scale = prompt.get("scale", 7.5)
|
||||
seed = prompt.get("seed")
|
||||
prompt = prompt.get("prompt")
|
||||
|
||||
prompt = replace_filewords_prompt(prompt, args)
|
||||
negative_prompt = replace_filewords_prompt(negative_prompt, args)
|
||||
else:
|
||||
prompt = replace_filewords_prompt(prompt, args)
|
||||
# prompt = prompt.strip()
|
||||
# if len(prompt) == 0 or prompt[0] == "#":
|
||||
# continue
|
||||
|
||||
# subset of gen_img_diffusers
|
||||
prompt_args = prompt.split(" --")
|
||||
prompt = prompt_args[0]
|
||||
negative_prompt = None
|
||||
sample_steps = 30
|
||||
width = height = 512
|
||||
scale = 7.5
|
||||
seed = None
|
||||
for parg in prompt_args:
|
||||
try:
|
||||
m = re.match(r"w (\d+)", parg, re.IGNORECASE)
|
||||
if m:
|
||||
width = int(m.group(1))
|
||||
continue
|
||||
|
||||
m = re.match(r"h (\d+)", parg, re.IGNORECASE)
|
||||
if m:
|
||||
height = int(m.group(1))
|
||||
continue
|
||||
|
||||
m = re.match(r"d (\d+)", parg, re.IGNORECASE)
|
||||
if m:
|
||||
seed = int(m.group(1))
|
||||
continue
|
||||
|
||||
m = re.match(r"s (\d+)", parg, re.IGNORECASE)
|
||||
if m: # steps
|
||||
sample_steps = max(1, min(1000, int(m.group(1))))
|
||||
continue
|
||||
|
||||
m = re.match(r"l ([\d\.]+)", parg, re.IGNORECASE)
|
||||
if m: # scale
|
||||
scale = float(m.group(1))
|
||||
continue
|
||||
|
||||
m = re.match(r"n (.+)", parg, re.IGNORECASE)
|
||||
if m: # negative prompt
|
||||
negative_prompt = m.group(1)
|
||||
continue
|
||||
|
||||
except ValueError as ex:
|
||||
print(f"Exception in parsing / 解析エラー: {parg}")
|
||||
print(ex)
|
||||
|
||||
if seed is not None:
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
if prompt_replacement is not None:
|
||||
prompt = prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||
if negative_prompt is not None:
|
||||
negative_prompt = negative_prompt.replace(prompt_replacement[0], prompt_replacement[1])
|
||||
|
||||
height = max(64, height - height % 8) # round to divisible by 8
|
||||
width = max(64, width - width % 8) # round to divisible by 8
|
||||
print(f"prompt: {prompt}")
|
||||
print(f"negative_prompt: {negative_prompt}")
|
||||
print(f"height: {height}")
|
||||
print(f"width: {width}")
|
||||
print(f"sample_steps: {sample_steps}")
|
||||
print(f"scale: {scale}")
|
||||
image = pipeline(
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_inference_steps=sample_steps,
|
||||
guidance_scale=scale,
|
||||
negative_prompt=negative_prompt,
|
||||
).images[0]
|
||||
|
||||
ts_str = time.strftime("%Y%m%d%H%M%S", time.localtime())
|
||||
num_suffix = f"e{epoch:06d}" if epoch is not None else f"{steps:06d}"
|
||||
seed_suffix = "" if seed is None else f"_{seed}"
|
||||
|
||||
if is_generating_only:
|
||||
img_filename = (
|
||||
f"{'' if args.output_name is None else args.output_name + '_'}{ts_str}_{num_suffix}_{i:02d}{seed_suffix}.png"
|
||||
)
|
||||
else:
|
||||
img_filename = (
|
||||
f"{'' if args.output_name is None else args.output_name + '_'}{ts_str}_{i:04d}{seed_suffix}.png"
|
||||
)
|
||||
if is_sample_only:
|
||||
# make prompt txt file
|
||||
img_path_no_ext = os.path.join(save_dir, img_filename[:-4])
|
||||
with open(img_path_no_ext + ".txt", "w") as f:
|
||||
# put prompt in txt file
|
||||
f.write(prompt)
|
||||
# close file
|
||||
f.close()
|
||||
|
||||
image.save(os.path.join(save_dir, img_filename))
|
||||
|
||||
# wandb有効時のみログを送信
|
||||
try:
|
||||
wandb_tracker = accelerator.get_tracker("wandb")
|
||||
try:
|
||||
import wandb
|
||||
except ImportError: # 事前に一度確認するのでここはエラー出ないはず
|
||||
raise ImportError("No wandb / wandb がインストールされていないようです")
|
||||
|
||||
wandb_tracker.log({f"sample_{i}": wandb.Image(image)})
|
||||
except: # wandb 無効時
|
||||
pass
|
||||
|
||||
# clear pipeline and cache to reduce vram usage
|
||||
del pipeline
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
torch.set_rng_state(rng_state)
|
||||
if cuda_rng_state is not None:
|
||||
torch.cuda.set_rng_state(cuda_rng_state)
|
||||
vae.to(org_vae_device)
|
||||
|
||||
|
||||
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
|
||||
def apply_noise_offset(noise, noise_offset):
|
||||
if noise_offset is None or noise_offset < 0.0000001:
|
||||
if noise_offset is None or (noise_offset < 0.000001 and noise_offset > -0.000001):
|
||||
return noise
|
||||
noise = noise + noise_offset * torch.randn((noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
|
||||
return noise
|
||||
@@ -579,6 +325,58 @@ def encode_prompts_xl(
|
||||
|
||||
return torch.concat(text_embeds_list, dim=-1), pooled_text_embeds
|
||||
|
||||
def encode_prompts_sd3(
|
||||
tokenizers: list['CLIPTokenizer'],
|
||||
text_encoders: list[Union['CLIPTextModel', 'CLIPTextModelWithProjection', T5EncoderModel]],
|
||||
prompts: list[str],
|
||||
num_images_per_prompt: int = 1,
|
||||
truncate: bool = True,
|
||||
max_length=None,
|
||||
dropout_prob=0.0,
|
||||
pipeline = None,
|
||||
):
|
||||
text_embeds_list = []
|
||||
pooled_text_embeds = None # always text_encoder_2's pool
|
||||
|
||||
prompt_2 = prompts
|
||||
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
|
||||
|
||||
prompt_3 = prompts
|
||||
prompt_3 = [prompt_3] if isinstance(prompt_3, str) else prompt_3
|
||||
|
||||
device = text_encoders[0].device
|
||||
|
||||
prompt_embed, pooled_prompt_embed = pipeline._get_clip_prompt_embeds(
|
||||
prompt=prompts,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
clip_skip=None,
|
||||
clip_model_index=0,
|
||||
)
|
||||
prompt_2_embed, pooled_prompt_2_embed = pipeline._get_clip_prompt_embeds(
|
||||
prompt=prompt_2,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
clip_skip=None,
|
||||
clip_model_index=1,
|
||||
)
|
||||
clip_prompt_embeds = torch.cat([prompt_embed, prompt_2_embed], dim=-1)
|
||||
|
||||
t5_prompt_embed = pipeline._get_t5_prompt_embeds(
|
||||
prompt=prompt_3,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device
|
||||
)
|
||||
|
||||
clip_prompt_embeds = torch.nn.functional.pad(
|
||||
clip_prompt_embeds, (0, t5_prompt_embed.shape[-1] - clip_prompt_embeds.shape[-1])
|
||||
)
|
||||
|
||||
prompt_embeds = torch.cat([clip_prompt_embeds, t5_prompt_embed], dim=-2)
|
||||
pooled_prompt_embeds = torch.cat([pooled_prompt_embed, pooled_prompt_2_embed], dim=-1)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
|
||||
# ref for long prompts https://github.com/huggingface/diffusers/issues/2136
|
||||
def text_encode(text_encoder: 'CLIPTextModel', tokens, truncate: bool = True, max_length=None):
|
||||
@@ -627,6 +425,159 @@ def encode_prompts(
|
||||
return text_embeddings
|
||||
|
||||
|
||||
def encode_prompts_pixart(
|
||||
tokenizer: 'T5Tokenizer',
|
||||
text_encoder: 'T5EncoderModel',
|
||||
prompts: list[str],
|
||||
truncate: bool = True,
|
||||
max_length=None,
|
||||
dropout_prob=0.0,
|
||||
):
|
||||
if max_length is None:
|
||||
# See Section 3.1. of the paper.
|
||||
max_length = 120
|
||||
|
||||
if dropout_prob > 0.0:
|
||||
# randomly drop out prompts
|
||||
prompts = [
|
||||
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
|
||||
]
|
||||
|
||||
text_inputs = tokenizer(
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = tokenizer(prompts, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
|
||||
text_input_ids, untruncated_ids
|
||||
):
|
||||
removed_text = tokenizer.batch_decode(untruncated_ids[:, max_length - 1: -1])
|
||||
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.to(text_encoder.device)
|
||||
|
||||
text_input_ids = text_input_ids.to(text_encoder.device)
|
||||
|
||||
prompt_embeds = text_encoder(text_input_ids, attention_mask=prompt_attention_mask)
|
||||
|
||||
return prompt_embeds.last_hidden_state, prompt_attention_mask
|
||||
|
||||
|
||||
def encode_prompts_auraflow(
|
||||
tokenizer: 'T5Tokenizer',
|
||||
text_encoder: 'UMT5EncoderModel',
|
||||
prompts: list[str],
|
||||
truncate: bool = True,
|
||||
max_length=None,
|
||||
dropout_prob=0.0,
|
||||
):
|
||||
if max_length is None:
|
||||
max_length = 256
|
||||
|
||||
if dropout_prob > 0.0:
|
||||
# randomly drop out prompts
|
||||
prompts = [
|
||||
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
|
||||
]
|
||||
|
||||
device = text_encoder.device
|
||||
|
||||
text_inputs = tokenizer(
|
||||
prompts,
|
||||
truncation=True,
|
||||
max_length=max_length,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs["input_ids"]
|
||||
untruncated_ids = tokenizer(prompts, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(
|
||||
text_input_ids, untruncated_ids
|
||||
):
|
||||
removed_text = tokenizer.batch_decode(untruncated_ids[:, max_length - 1: -1])
|
||||
|
||||
text_inputs = {k: v.to(device) for k, v in text_inputs.items()}
|
||||
prompt_embeds = text_encoder(**text_inputs)[0]
|
||||
prompt_attention_mask = text_inputs["attention_mask"].unsqueeze(-1).expand(prompt_embeds.shape)
|
||||
prompt_embeds = prompt_embeds * prompt_attention_mask
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def encode_prompts_flux(
|
||||
tokenizer: List[Union['CLIPTokenizer','T5Tokenizer']],
|
||||
text_encoder: List[Union['CLIPTextModel', 'T5EncoderModel']],
|
||||
prompts: list[str],
|
||||
truncate: bool = True,
|
||||
max_length=None,
|
||||
dropout_prob=0.0,
|
||||
):
|
||||
if max_length is None:
|
||||
max_length = 512
|
||||
|
||||
if dropout_prob > 0.0:
|
||||
# randomly drop out prompts
|
||||
prompts = [
|
||||
prompt if torch.rand(1).item() > dropout_prob else "" for prompt in prompts
|
||||
]
|
||||
|
||||
device = text_encoder[0].device
|
||||
dtype = text_encoder[0].dtype
|
||||
|
||||
batch_size = len(prompts)
|
||||
|
||||
# clip
|
||||
text_inputs = tokenizer[0](
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=tokenizer[0].model_max_length,
|
||||
truncation=True,
|
||||
return_overflowing_tokens=False,
|
||||
return_length=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids
|
||||
|
||||
prompt_embeds = text_encoder[0](text_input_ids.to(device), output_hidden_states=False)
|
||||
|
||||
# Use pooled output of CLIPTextModel
|
||||
pooled_prompt_embeds = prompt_embeds.pooler_output
|
||||
pooled_prompt_embeds = pooled_prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# T5
|
||||
text_inputs = tokenizer[1](
|
||||
prompts,
|
||||
padding="max_length",
|
||||
max_length=max_length,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
|
||||
prompt_embeds = text_encoder[1](text_input_ids.to(device), output_hidden_states=False)[0]
|
||||
|
||||
dtype = text_encoder[1].dtype
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# prompt_attention_mask = text_inputs["attention_mask"].unsqueeze(-1).expand(prompt_embeds.shape)
|
||||
# prompt_embeds = prompt_embeds * prompt_attention_mask
|
||||
# _, seq_len, _ = prompt_embeds.shape
|
||||
|
||||
# they dont do prompt attention mask?
|
||||
# prompt_attention_mask = torch.ones((batch_size, seq_len), dtype=dtype, device=device)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
|
||||
# for XL
|
||||
def get_add_time_ids(
|
||||
height: int,
|
||||
@@ -675,18 +626,22 @@ def concat_embeddings(
|
||||
|
||||
|
||||
def add_all_snr_to_noise_scheduler(noise_scheduler, device):
|
||||
if hasattr(noise_scheduler, "all_snr"):
|
||||
return
|
||||
# compute it
|
||||
with torch.no_grad():
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
alpha = sqrt_alphas_cumprod
|
||||
sigma = sqrt_one_minus_alphas_cumprod
|
||||
all_snr = (alpha / sigma) ** 2
|
||||
all_snr.requires_grad = False
|
||||
noise_scheduler.all_snr = all_snr.to(device)
|
||||
try:
|
||||
if hasattr(noise_scheduler, "all_snr"):
|
||||
return
|
||||
# compute it
|
||||
with torch.no_grad():
|
||||
alphas_cumprod = noise_scheduler.alphas_cumprod
|
||||
sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
alpha = sqrt_alphas_cumprod
|
||||
sigma = sqrt_one_minus_alphas_cumprod
|
||||
all_snr = (alpha / sigma) ** 2
|
||||
all_snr.requires_grad = False
|
||||
noise_scheduler.all_snr = all_snr.to(device)
|
||||
except Exception as e:
|
||||
# just move on
|
||||
pass
|
||||
|
||||
|
||||
def get_all_snr(noise_scheduler, device):
|
||||
@@ -776,8 +731,19 @@ def apply_snr_weight(
|
||||
):
|
||||
# will get it from noise scheduler if exist or will calculate it if not
|
||||
all_snr = get_all_snr(noise_scheduler, loss.device)
|
||||
step_indices = [(noise_scheduler.timesteps == t).nonzero().item() for t in timesteps]
|
||||
snr = torch.stack([all_snr[t] for t in step_indices])
|
||||
# step_indices = []
|
||||
# for t in timesteps:
|
||||
# for i, st in enumerate(noise_scheduler.timesteps):
|
||||
# if st == t:
|
||||
# step_indices.append(i)
|
||||
# break
|
||||
# this breaks on some schedulers
|
||||
# step_indices = [(noise_scheduler.timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
offset = 0
|
||||
if noise_scheduler.timesteps[0] == 1000:
|
||||
offset = 1
|
||||
snr = torch.stack([all_snr[(t - offset).int()] for t in timesteps])
|
||||
gamma_over_snr = torch.div(torch.ones_like(snr) * gamma, snr)
|
||||
if fixed:
|
||||
snr_weight = gamma_over_snr.float().to(loss.device) # directly using gamma over snr
|
||||
@@ -786,3 +752,19 @@ def apply_snr_weight(
|
||||
snr_adjusted_loss = loss * snr_weight
|
||||
|
||||
return snr_adjusted_loss
|
||||
|
||||
|
||||
def precondition_model_outputs_flow_match(model_output, model_input, timestep_tensor, noise_scheduler):
|
||||
mo_chunks = torch.chunk(model_output, model_output.shape[0], dim=0)
|
||||
mi_chunks = torch.chunk(model_input, model_input.shape[0], dim=0)
|
||||
timestep_chunks = torch.chunk(timestep_tensor, timestep_tensor.shape[0], dim=0)
|
||||
out_chunks = []
|
||||
# unsqueeze if timestep is zero dim
|
||||
for idx in range(model_output.shape[0]):
|
||||
sigmas = noise_scheduler.get_sigmas(timestep_chunks[idx], n_dim=model_output.ndim,
|
||||
dtype=model_output.dtype, device=model_output.device)
|
||||
# Follow: Section 5 of https://arxiv.org/abs/2206.00364.
|
||||
# Preconditioning of the model outputs.
|
||||
out = mo_chunks[idx] * (-sigmas) + mi_chunks[idx]
|
||||
out_chunks.append(out)
|
||||
return torch.cat(out_chunks, dim=0)
|
||||
|
||||
120
toolkit/util/adafactor_stochastic_rounding.py
Normal file
120
toolkit/util/adafactor_stochastic_rounding.py
Normal file
@@ -0,0 +1,120 @@
|
||||
# ref https://github.com/Nerogar/OneTrainer/compare/master...stochastic_rounding
|
||||
import math
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def copy_stochastic_(target: Tensor, source: Tensor):
|
||||
# create a random 16 bit integer
|
||||
result = torch.randint_like(
|
||||
source,
|
||||
dtype=torch.int32,
|
||||
low=0,
|
||||
high=(1 << 16),
|
||||
)
|
||||
|
||||
# add the random number to the lower 16 bit of the mantissa
|
||||
result.add_(source.view(dtype=torch.int32))
|
||||
|
||||
# mask off the lower 16 bit of the mantissa
|
||||
result.bitwise_and_(-65536) # -65536 = FFFF0000 as a signed int32
|
||||
|
||||
# copy the higher 16 bit into the target tensor
|
||||
target.copy_(result.view(dtype=torch.float32))
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def step_adafactor(self, closure=None):
|
||||
"""
|
||||
Performs a single optimization step
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
for p in group["params"]:
|
||||
if p.grad is None:
|
||||
continue
|
||||
grad = p.grad
|
||||
if grad.dtype in {torch.float16, torch.bfloat16}:
|
||||
grad = grad.float()
|
||||
if grad.is_sparse:
|
||||
raise RuntimeError("Adafactor does not support sparse gradients.")
|
||||
|
||||
state = self.state[p]
|
||||
grad_shape = grad.shape
|
||||
|
||||
factored, use_first_moment = self._get_options(group, grad_shape)
|
||||
# State Initialization
|
||||
if len(state) == 0:
|
||||
state["step"] = 0
|
||||
|
||||
if use_first_moment:
|
||||
# Exponential moving average of gradient values
|
||||
state["exp_avg"] = torch.zeros_like(grad)
|
||||
if factored:
|
||||
state["exp_avg_sq_row"] = torch.zeros(grad_shape[:-1]).to(grad)
|
||||
state["exp_avg_sq_col"] = torch.zeros(grad_shape[:-2] + grad_shape[-1:]).to(grad)
|
||||
else:
|
||||
state["exp_avg_sq"] = torch.zeros_like(grad)
|
||||
|
||||
state["RMS"] = 0
|
||||
else:
|
||||
if use_first_moment:
|
||||
state["exp_avg"] = state["exp_avg"].to(grad)
|
||||
if factored:
|
||||
state["exp_avg_sq_row"] = state["exp_avg_sq_row"].to(grad)
|
||||
state["exp_avg_sq_col"] = state["exp_avg_sq_col"].to(grad)
|
||||
else:
|
||||
state["exp_avg_sq"] = state["exp_avg_sq"].to(grad)
|
||||
|
||||
p_data_fp32 = p
|
||||
if p.dtype in {torch.float16, torch.bfloat16}:
|
||||
p_data_fp32 = p_data_fp32.float()
|
||||
|
||||
state["step"] += 1
|
||||
state["RMS"] = self._rms(p_data_fp32)
|
||||
lr = self._get_lr(group, state)
|
||||
|
||||
beta2t = 1.0 - math.pow(state["step"], group["decay_rate"])
|
||||
eps = group["eps"][0] if isinstance(group["eps"], list) else group["eps"]
|
||||
update = (grad ** 2) + eps
|
||||
if factored:
|
||||
exp_avg_sq_row = state["exp_avg_sq_row"]
|
||||
exp_avg_sq_col = state["exp_avg_sq_col"]
|
||||
|
||||
exp_avg_sq_row.mul_(beta2t).add_(update.mean(dim=-1), alpha=(1.0 - beta2t))
|
||||
exp_avg_sq_col.mul_(beta2t).add_(update.mean(dim=-2), alpha=(1.0 - beta2t))
|
||||
|
||||
# Approximation of exponential moving average of square of gradient
|
||||
update = self._approx_sq_grad(exp_avg_sq_row, exp_avg_sq_col)
|
||||
update.mul_(grad)
|
||||
else:
|
||||
exp_avg_sq = state["exp_avg_sq"]
|
||||
|
||||
exp_avg_sq.mul_(beta2t).add_(update, alpha=(1.0 - beta2t))
|
||||
update = exp_avg_sq.rsqrt().mul_(grad)
|
||||
|
||||
update.div_((self._rms(update) / group["clip_threshold"]).clamp_(min=1.0))
|
||||
update.mul_(lr)
|
||||
|
||||
if use_first_moment:
|
||||
exp_avg = state["exp_avg"]
|
||||
exp_avg.mul_(group["beta1"]).add_(update, alpha=(1 - group["beta1"]))
|
||||
update = exp_avg
|
||||
|
||||
if group["weight_decay"] != 0:
|
||||
p_data_fp32.add_(p_data_fp32, alpha=(-group["weight_decay"] * lr))
|
||||
|
||||
p_data_fp32.add_(-update)
|
||||
|
||||
if p.dtype == torch.bfloat16:
|
||||
copy_stochastic_(p, p_data_fp32)
|
||||
elif p.dtype == torch.float16:
|
||||
p.copy_(p_data_fp32)
|
||||
|
||||
return loss
|
||||
25
toolkit/util/inverse_cfg.py
Normal file
25
toolkit/util/inverse_cfg.py
Normal file
@@ -0,0 +1,25 @@
|
||||
import torch
|
||||
|
||||
|
||||
def inverse_classifier_guidance(
|
||||
noise_pred_cond: torch.Tensor,
|
||||
noise_pred_uncond: torch.Tensor,
|
||||
guidance_scale: torch.Tensor
|
||||
):
|
||||
"""
|
||||
Adjust the noise_pred_cond for the classifier free guidance algorithm
|
||||
to ensure that the final noise prediction equals the original noise_pred_cond.
|
||||
"""
|
||||
# To make noise_pred equal noise_pred_cond_orig, we adjust noise_pred_cond
|
||||
# based on the formula used in the algorithm.
|
||||
# We derive the formula to find the correct adjustment for noise_pred_cond:
|
||||
# noise_pred_cond = (noise_pred_cond_orig - noise_pred_uncond * guidance_scale) / (guidance_scale - 1)
|
||||
# It's important to check if guidance_scale is not 1 to avoid division by zero.
|
||||
if guidance_scale == 1:
|
||||
# If guidance_scale is 1, adjusting is not needed or possible in the same way,
|
||||
# since it would lead to division by zero. This also means the algorithm inherently
|
||||
# doesn't alter the noise_pred_cond in relation to noise_pred_uncond.
|
||||
# Thus, we return the original values, though this situation might need special handling.
|
||||
return noise_pred_cond
|
||||
adjusted_noise_pred_cond = (noise_pred_cond - noise_pred_uncond) / guidance_scale
|
||||
return adjusted_noise_pred_cond
|
||||
Reference in New Issue
Block a user