Compare commits
211 Commits
60232def91
...
wavelet_lo
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ce4c5291a0 | ||
|
|
c101f07834 | ||
|
|
e4526ad4a4 | ||
|
|
4595965e06 | ||
|
|
41edc18750 | ||
|
|
6021a3dbc0 | ||
|
|
71d7a52146 | ||
|
|
45be82d5d6 | ||
|
|
f10937e6da | ||
|
|
ccb66c748f | ||
|
|
2aca2883e7 | ||
|
|
1ad58c5816 | ||
|
|
6dea41b9fc | ||
|
|
9a902c067f | ||
|
|
0bbc69c135 | ||
|
|
6c5eb0cf87 | ||
|
|
e3373671b9 | ||
|
|
aceb3a0f25 | ||
|
|
c8049a483d | ||
|
|
f5aa4232fa | ||
|
|
3a6b24f4c8 | ||
|
|
bbfd6ef0fe | ||
|
|
b829983b16 | ||
|
|
fa187b1208 | ||
|
|
5eb627dd9d | ||
|
|
604e76d34d | ||
|
|
6cde96ae5f | ||
|
|
1be613ed06 | ||
|
|
c52421aab7 | ||
|
|
3812957bc9 | ||
|
|
391329dbdc | ||
|
|
3b45892b4f | ||
|
|
cf4216e6b8 | ||
|
|
31e057d9a3 | ||
|
|
d507b44a7b | ||
|
|
242c04a0b8 | ||
|
|
386e68a422 | ||
|
|
850b8da6e5 | ||
|
|
51ad19b568 | ||
|
|
e6739f7eb2 | ||
|
|
7e37918fbc | ||
|
|
4d88f8f218 | ||
|
|
25341c4613 | ||
|
|
391cf80fea | ||
|
|
4e3bda7c70 | ||
|
|
763128ea42 | ||
|
|
4fe33f51c1 | ||
|
|
aa44828c0c | ||
|
|
6f6fb90812 | ||
|
|
c57434ad7b | ||
|
|
8bb47d1bfe | ||
|
|
e7dbb20f68 | ||
|
|
c5e0c2bbe2 | ||
|
|
1f3f45a48d | ||
|
|
3c8c84f156 | ||
|
|
b001d77efb | ||
|
|
7ae31c9ae9 | ||
|
|
b16819f8e7 | ||
|
|
f5e40dfa62 | ||
|
|
acc79956aa | ||
|
|
60539c0b0f | ||
|
|
dd700f70b3 | ||
|
|
d360e76661 | ||
|
|
6ec23ed226 | ||
|
|
f6e16e582a | ||
|
|
259ded9602 | ||
|
|
440ba5fb3d | ||
|
|
093f14ac19 | ||
|
|
f0fbd8bb53 | ||
|
|
0a981bea2b | ||
|
|
1d0e3a4498 | ||
|
|
3c7daf49f3 | ||
|
|
56d8d6bd81 | ||
|
|
3e49337a58 | ||
|
|
60f848a877 | ||
|
|
b366e46f1c | ||
|
|
a280f78c69 | ||
|
|
6e19e7449e | ||
|
|
a6d46ad9ae | ||
|
|
f3725578dd | ||
|
|
ed99c3c0c8 | ||
|
|
ed84c19205 | ||
|
|
a7a9c11d9e | ||
|
|
f60698d0ee | ||
|
|
5f094fb17a | ||
|
|
a5227cba7b | ||
|
|
77a5e01301 | ||
|
|
4ef5a668c0 | ||
|
|
f081d14527 | ||
|
|
710c6de1c9 | ||
|
|
2b6e66e0cb | ||
|
|
ab641e014f | ||
|
|
ad87f72384 | ||
|
|
d0214c0df9 | ||
|
|
adcf884c0f | ||
|
|
f778d979b5 | ||
|
|
db3ccbba33 | ||
|
|
0d2be18a9b | ||
|
|
bbc340e545 | ||
|
|
33fdfd6091 | ||
|
|
9f6030620f | ||
|
|
b5252b5028 | ||
|
|
b0d8fc220d | ||
|
|
cef7d9e594 | ||
|
|
b13fcc1039 | ||
|
|
b32d7e552b | ||
|
|
4af6c5cf30 | ||
|
|
1f7784510d | ||
|
|
87e557cf1e | ||
|
|
bd8d7dc081 | ||
|
|
2be6926398 | ||
|
|
87ac031859 | ||
|
|
7679105d52 | ||
|
|
2622de1e01 | ||
|
|
8450aca10e | ||
|
|
0b8a32def7 | ||
|
|
787bb37e76 | ||
|
|
10aa7e9d5e | ||
|
|
ed1deb71c4 | ||
|
|
4de6a825fa | ||
|
|
9a7266275d | ||
|
|
d138f07365 | ||
|
|
c6d8eedb94 | ||
|
|
af5e760be1 | ||
|
|
ff3d54bb5b | ||
|
|
0e75724b4d | ||
|
|
376bb1bf6f | ||
|
|
216ab164ce | ||
|
|
e6180d1e1d | ||
|
|
15a57bc89f | ||
|
|
e5355bf8d5 | ||
|
|
34a1c6947a | ||
|
|
2141c6e06c | ||
|
|
1188cf1e8a | ||
|
|
5e663746b8 | ||
|
|
441474e81f | ||
|
|
a6a690f796 | ||
|
|
6191f19e55 | ||
|
|
bbfba0c188 | ||
|
|
e1549ad54d | ||
|
|
04abe57c76 | ||
|
|
89dd041b97 | ||
|
|
29122b1a54 | ||
|
|
6a8e3d8610 | ||
|
|
4c8a9e1b88 | ||
|
|
fadb2f3a76 | ||
|
|
4723f23c0d | ||
|
|
8ef07a9c36 | ||
|
|
92ce93140e | ||
|
|
f213996aa5 | ||
|
|
cbe31eaf0a | ||
|
|
67c2e44edb | ||
|
|
96d418bb95 | ||
|
|
894374b2e9 | ||
|
|
6509ba4484 | ||
|
|
025ee3dd3d | ||
|
|
58f9d01c2b | ||
|
|
e72b59a8e9 | ||
|
|
4aa19b5c1d | ||
|
|
4747716867 | ||
|
|
22cd40d7b9 | ||
|
|
3400882a80 | ||
|
|
9f94c7b61e | ||
|
|
bedb8197a2 | ||
|
|
e3ebd73610 | ||
|
|
dd931757cd | ||
|
|
0640cdf569 | ||
|
|
0b048d0dde | ||
|
|
473d455f44 | ||
|
|
ce759ebd8c | ||
|
|
628a7923a3 | ||
|
|
3922981996 | ||
|
|
ab22674980 | ||
|
|
9452929300 | ||
|
|
a800c9d19e | ||
|
|
28e6f00790 | ||
|
|
67e0aca750 | ||
|
|
f05224970f | ||
|
|
b4f64de4c2 | ||
|
|
2e5f6668dc | ||
|
|
e4c82803e1 | ||
|
|
69aa92bce5 | ||
|
|
a508caad1d | ||
|
|
58537fc92b | ||
|
|
86b5938cf3 | ||
|
|
6b4034122f | ||
|
|
10817696fb | ||
|
|
037ce11740 | ||
|
|
04424fe2d6 | ||
|
|
40a8ff5731 | ||
|
|
2776221497 | ||
|
|
f85ad452c6 | ||
|
|
dd889086f4 | ||
|
|
bc693488eb | ||
|
|
d97c55cd96 | ||
|
|
79b4e04b80 | ||
|
|
951e223481 | ||
|
|
fc34a69bec | ||
|
|
279ee65177 | ||
|
|
3a1f464132 | ||
|
|
5c8fcc8a4e | ||
|
|
121a760c19 | ||
|
|
e5fadddd45 | ||
|
|
d44d4eb61a | ||
|
|
7d9ab22405 | ||
|
|
7ed8c51f20 | ||
|
|
6df33156f0 | ||
|
|
40f5c59da0 | ||
|
|
3e71a99df0 | ||
|
|
562405923f | ||
|
|
f84bd6d7a6 |
2
.github/FUNDING.yml
vendored
Normal file
2
.github/FUNDING.yml
vendored
Normal file
@@ -0,0 +1,2 @@
|
||||
github: [ostris]
|
||||
patreon: ostris
|
||||
1
.github/ISSUE_TEMPLATE/bug_report.md
vendored
1
.github/ISSUE_TEMPLATE/bug_report.md
vendored
@@ -17,4 +17,3 @@ You verified that this is a bug and not a feature request or question by asking
|
||||
Yes/No
|
||||
|
||||
## Describe the bug
|
||||
|
||||
|
||||
8
.gitignore
vendored
8
.gitignore
vendored
@@ -161,6 +161,7 @@ cython_debug/
|
||||
|
||||
/env.sh
|
||||
/models
|
||||
/datasets
|
||||
/custom/*
|
||||
!/custom/.gitkeep
|
||||
/.tmp
|
||||
@@ -173,4 +174,9 @@ cython_debug/
|
||||
!/output/.gitkeep
|
||||
/extensions/*
|
||||
!/extensions/example
|
||||
/temp
|
||||
/temp
|
||||
/wandb
|
||||
.vscode/settings.json
|
||||
.DS_Store
|
||||
._.DS_Store
|
||||
aitk_db.db
|
||||
4
.gitmodules
vendored
4
.gitmodules
vendored
@@ -1,12 +1,16 @@
|
||||
[submodule "repositories/sd-scripts"]
|
||||
path = repositories/sd-scripts
|
||||
url = https://github.com/kohya-ss/sd-scripts.git
|
||||
commit = b78c0e2a69e52ce6c79abc6c8c82d1a9cabcf05c
|
||||
[submodule "repositories/leco"]
|
||||
path = repositories/leco
|
||||
url = https://github.com/p1atdev/LECO
|
||||
commit = 9294adf40218e917df4516737afb13f069a6789d
|
||||
[submodule "repositories/batch_annotator"]
|
||||
path = repositories/batch_annotator
|
||||
url = https://github.com/ostris/batch-annotator
|
||||
commit = 420e142f6ad3cc14b3ea0500affc2c6c7e7544bf
|
||||
[submodule "repositories/ipadapter"]
|
||||
path = repositories/ipadapter
|
||||
url = https://github.com/tencent-ailab/IP-Adapter.git
|
||||
commit = 5a18b1f3660acaf8bee8250692d6fb3548a19b14
|
||||
|
||||
28
.vscode/launch.json
vendored
Normal file
28
.vscode/launch.json
vendored
Normal file
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "Run current config",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/run.py",
|
||||
"args": [
|
||||
"${file}"
|
||||
],
|
||||
"env": {
|
||||
"CUDA_LAUNCH_BLOCKING": "1",
|
||||
"DEBUG_TOOLKIT": "1"
|
||||
},
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "Python: Debug Current File",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
]
|
||||
}
|
||||
BIN
assets/lora_ease_ui.png
Normal file
BIN
assets/lora_ease_ui.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 340 KiB |
29
build_and_push_docker
Normal file
29
build_and_push_docker
Normal file
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
# Extract version from version.py
|
||||
if [ -f "version.py" ]; then
|
||||
VERSION=$(python3 -c "from version import VERSION; print(VERSION)")
|
||||
echo "Building version: $VERSION"
|
||||
else
|
||||
echo "Error: version.py not found. Please create a version.py file with VERSION defined."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
|
||||
echo "Building version: $VERSION and latest"
|
||||
# wait 2 seconds
|
||||
sleep 2
|
||||
|
||||
# Build the image with cache busting
|
||||
docker build --build-arg CACHEBUST=$(date +%s) -t aitoolkit:$VERSION -f docker/Dockerfile .
|
||||
|
||||
# Tag with version and latest
|
||||
docker tag aitoolkit:$VERSION ostris/aitoolkit:$VERSION
|
||||
docker tag aitoolkit:$VERSION ostris/aitoolkit:latest
|
||||
|
||||
# Push both tags
|
||||
echo "Pushing images to Docker Hub..."
|
||||
docker push ostris/aitoolkit:$VERSION
|
||||
docker push ostris/aitoolkit:latest
|
||||
|
||||
echo "Successfully built and pushed ostris/aitoolkit:$VERSION and ostris/aitoolkit:latest"
|
||||
@@ -1,8 +0,0 @@
|
||||
#!/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
|
||||
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
@@ -0,0 +1,107 @@
|
||||
---
|
||||
# This configuration requires 48GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_finetune_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"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
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 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lognorm_blend'
|
||||
timestep_type: 'sigmoid'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true # flex is flux architecture
|
||||
# full finetuning quantized models is a crapshoot and results in subpar outputs
|
||||
# quantize: true
|
||||
# you can quantize just the T5 text encoder here to save vram
|
||||
quantize_te: true
|
||||
# only train the transformer blocks
|
||||
only_if_contains:
|
||||
- "transformer.transformer_blocks."
|
||||
- "transformer.single_transformer_blocks."
|
||||
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 flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
@@ -0,0 +1,99 @@
|
||||
---
|
||||
# This configuration requires 24GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_finetune_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"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
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 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
# ema_config:
|
||||
# use_ema: true
|
||||
# ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
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 cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "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: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
101
config/examples/train_lora_flex_24gb.yaml
Normal file
101
config/examples/train_lora_flex_24gb.yaml
Normal file
@@ -0,0 +1,101 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_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 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
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 flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
quantize_kwargs:
|
||||
exclude:
|
||||
- "*time_text_embed*" # exclude the time text embedder from quantization
|
||||
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 flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# 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_lumina.yaml
Normal file
96
config/examples/train_lora_lumina.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
# This configuration requires 20GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_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: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
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 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
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
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
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 cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "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: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
@@ -0,0 +1,97 @@
|
||||
---
|
||||
# NOTE!! THIS IS CURRENTLY EXPERIMENTAL AND UNDER DEVELOPMENT. SOME THINGS WILL CHANGE
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_sd3l_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: [ 1024 ]
|
||||
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 # May not fully work with SD3 yet
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch"
|
||||
timestep_type: "linear" # linear or sigmoid
|
||||
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 sd3, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "stabilityai/stable-diffusion-3.5-large"
|
||||
is_v3: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
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: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
101
config/examples/train_lora_wan21_14b_24gb.yaml
Normal file
101
config/examples/train_lora_wan21_14b_24gb.yaml
Normal file
@@ -0,0 +1,101 @@
|
||||
# IMPORTANT: The Wan2.1 14B model is huge. This config should work on 24GB GPUs. It cannot
|
||||
# support keeping the text encoder on GPU while training with 24GB, so it is only good
|
||||
# for training on a single prompt, for example a person with a trigger word.
|
||||
# to train on captions, you need more vran for now.
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_wan21_14b_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
|
||||
# this is probably needed for 24GB cards when offloading TE to CPU
|
||||
trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 32
|
||||
linear_alpha: 32
|
||||
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"
|
||||
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
|
||||
# it works well for characters, but not as well for "actions"
|
||||
- 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: [ 632 ] # will be around 480p
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with wan
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
timestep_type: 'sigmoid'
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
optimizer_params:
|
||||
weight_decay: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
dtype: bf16
|
||||
# required for 24GB cards
|
||||
# this will encode your trigger word and use those embeddings for every image in the dataset
|
||||
unload_text_encoder: true
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
arch: 'wan21'
|
||||
# these settings will save as much vram as possible
|
||||
quantize: true
|
||||
quantize_te: true
|
||||
low_vram: true
|
||||
sample:
|
||||
sampler: "flowmatch"
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 832
|
||||
height: 480
|
||||
num_frames: 40
|
||||
fps: 15
|
||||
# samples take a long time. so use them sparingly
|
||||
# samples will be animated webp files, if you don't see them animated, open in a browser.
|
||||
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 playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 5
|
||||
sample_steps: 30
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
90
config/examples/train_lora_wan21_1b_24gb.yaml
Normal file
90
config/examples/train_lora_wan21_1b_24gb.yaml
Normal file
@@ -0,0 +1,90 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_wan21_1b_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: 32
|
||||
linear_alpha: 32
|
||||
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"
|
||||
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
|
||||
# it works well for characters, but not as well for "actions"
|
||||
- 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: [ 632 ] # will be around 480p
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with wan
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
timestep_type: 'sigmoid'
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
optimizer_params:
|
||||
weight_decay: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
arch: 'wan21'
|
||||
quantize_te: true # saves vram
|
||||
sample:
|
||||
sampler: "flowmatch"
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 832
|
||||
height: 480
|
||||
num_frames: 40
|
||||
fps: 15
|
||||
# samples take a long time. so use them sparingly
|
||||
# samples will be animated webp files, if you don't see them animated, open in a browser.
|
||||
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 playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 5
|
||||
sample_steps: 30
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
25
docker-compose.yml
Normal file
25
docker-compose.yml
Normal file
@@ -0,0 +1,25 @@
|
||||
version: "3.8"
|
||||
|
||||
services:
|
||||
ai-toolkit:
|
||||
image: ostris/aitoolkit:latest
|
||||
restart: unless-stopped
|
||||
ports:
|
||||
- "8675:8675"
|
||||
volumes:
|
||||
- ~/.cache/huggingface/hub:/root/.cache/huggingface/hub
|
||||
- ./aitk_db.db:/app/ai-toolkit/aitk_db.db
|
||||
- ./datasets:/app/ai-toolkit/datasets
|
||||
- ./output:/app/ai-toolkit/output
|
||||
- ./config:/app/ai-toolkit/config
|
||||
environment:
|
||||
- AI_TOOLKIT_AUTH=${AI_TOOLKIT_AUTH:-password}
|
||||
- NODE_ENV=production
|
||||
- TZ=UTC
|
||||
deploy:
|
||||
resources:
|
||||
reservations:
|
||||
devices:
|
||||
- driver: nvidia
|
||||
count: all
|
||||
capabilities: [gpu]
|
||||
@@ -1,21 +1,67 @@
|
||||
FROM runpod/base:0.6.2-cuda12.1.0
|
||||
FROM nvidia/cuda:12.6.3-base-ubuntu22.04
|
||||
|
||||
LABEL authors="jaret"
|
||||
|
||||
# Set noninteractive to avoid timezone prompts
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
# Install dependencies
|
||||
RUN apt-get update
|
||||
RUN apt-get update && apt-get install --no-install-recommends -y \
|
||||
git \
|
||||
curl \
|
||||
build-essential \
|
||||
cmake \
|
||||
wget \
|
||||
python3.10 \
|
||||
python3-pip \
|
||||
python3-dev \
|
||||
python3-setuptools \
|
||||
python3-wheel \
|
||||
python3-venv \
|
||||
ffmpeg \
|
||||
tmux \
|
||||
htop \
|
||||
nvtop \
|
||||
python3-opencv \
|
||||
&& apt-get clean \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Install nodejs
|
||||
WORKDIR /tmp
|
||||
RUN curl -sL https://deb.nodesource.com/setup_23.x -o nodesource_setup.sh && \
|
||||
bash nodesource_setup.sh && \
|
||||
apt-get update && \
|
||||
apt-get install -y nodejs && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
WORKDIR /app
|
||||
ARG CACHEBUST=1
|
||||
RUN git clone https://github.com/ostris/ai-toolkit.git && \
|
||||
|
||||
# Set aliases for python and pip
|
||||
RUN ln -s /usr/bin/python3 /usr/bin/python
|
||||
|
||||
# install pytorch before cache bust to avoid redownloading pytorch
|
||||
RUN pip install --no-cache-dir torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/cu126
|
||||
|
||||
# Fix cache busting by moving CACHEBUST to right before git clone
|
||||
ARG CACHEBUST=1234
|
||||
RUN echo "Cache bust: ${CACHEBUST}" && \
|
||||
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
|
||||
# Install Python dependencies
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
RUN apt-get install -y tmux nvtop htop
|
||||
# Build UI
|
||||
WORKDIR /app/ai-toolkit/ui
|
||||
RUN npm install && \
|
||||
npm run build && \
|
||||
npm run update_db
|
||||
|
||||
WORKDIR /
|
||||
CMD ["/start.sh"]
|
||||
# Expose port (assuming the application runs on port 3000)
|
||||
EXPOSE 8675
|
||||
|
||||
CMD ["npm", "run", "start"]
|
||||
@@ -20,6 +20,7 @@ from toolkit.guidance import get_targeted_guidance_loss, get_guidance_loss, Guid
|
||||
from toolkit.image_utils import show_tensors, show_latents
|
||||
from toolkit.ip_adapter import IPAdapter
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
||||
from toolkit.reference_adapter import ReferenceAdapter
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, BlankNetwork
|
||||
@@ -32,6 +33,8 @@ from torchvision import transforms
|
||||
from diffusers import EMAModel
|
||||
import math
|
||||
from toolkit.train_tools import precondition_model_outputs_flow_match
|
||||
from toolkit.models.diffusion_feature_extraction import DiffusionFeatureExtractor, load_dfe
|
||||
from toolkit.util.wavelet_loss import wavelet_loss
|
||||
|
||||
|
||||
def flush():
|
||||
@@ -58,23 +61,38 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
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:
|
||||
if self.is_fine_tuning and self.is_bfloat:
|
||||
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
|
||||
# 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.cached_blank_embeds: Optional[PromptEmbeds] = None
|
||||
self.cached_trigger_embeds: Optional[PromptEmbeds] = None
|
||||
self.diff_output_preservation_embeds: Optional[PromptEmbeds] = None
|
||||
|
||||
self.dfe: Optional[DiffusionFeatureExtractor] = None
|
||||
|
||||
if self.train_config.diff_output_preservation:
|
||||
if self.trigger_word is None:
|
||||
raise ValueError("diff_output_preservation requires a trigger_word to be set")
|
||||
if self.network_config is None:
|
||||
raise ValueError("diff_output_preservation requires a network to be set")
|
||||
if self.train_config.train_text_encoder:
|
||||
raise ValueError("diff_output_preservation is not supported with train_text_encoder")
|
||||
|
||||
# always do a prior prediction when doing diff output preservation
|
||||
self.do_prior_prediction = True
|
||||
|
||||
|
||||
def before_model_load(self):
|
||||
@@ -113,6 +131,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
self.taesd.requires_grad_(False)
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
super().hook_before_train_loop()
|
||||
|
||||
if self.train_config.do_prior_divergence:
|
||||
self.do_prior_prediction = True
|
||||
# move vae to device if we did not cache latents
|
||||
@@ -153,6 +173,35 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
# single prompt
|
||||
self.negative_prompt_pool = [self.train_config.negative_prompt]
|
||||
|
||||
# handle unload text encoder
|
||||
if self.train_config.unload_text_encoder:
|
||||
with torch.no_grad():
|
||||
if self.train_config.train_text_encoder:
|
||||
raise ValueError("Cannot unload text encoder if training text encoder")
|
||||
# cache embeddings
|
||||
|
||||
print_acc("\n***** UNLOADING TEXT ENCODER *****")
|
||||
print_acc("This will train only with a blank prompt or trigger word, if set")
|
||||
print_acc("If this is not what you want, remove the unload_text_encoder flag")
|
||||
print_acc("***********************************")
|
||||
print_acc("")
|
||||
self.sd.text_encoder_to(self.device_torch)
|
||||
self.cached_blank_embeds = self.sd.encode_prompt("")
|
||||
if self.trigger_word is not None:
|
||||
self.cached_trigger_embeds = self.sd.encode_prompt(self.trigger_word)
|
||||
if self.train_config.diff_output_preservation:
|
||||
self.diff_output_preservation_embeds = self.sd.encode_prompt(self.train_config.diff_output_preservation_class)
|
||||
|
||||
# move back to cpu
|
||||
self.sd.text_encoder_to('cpu')
|
||||
flush()
|
||||
|
||||
if self.train_config.diffusion_feature_extractor_path is not None:
|
||||
self.dfe = load_dfe(self.train_config.diffusion_feature_extractor_path)
|
||||
self.dfe.to(self.device_torch)
|
||||
self.dfe.eval()
|
||||
|
||||
|
||||
def process_output_for_turbo(self, pred, noisy_latents, timesteps, noise, batch):
|
||||
# to process turbo learning, we make one big step from our current timestep to the end
|
||||
# we then denoise the prediction on that remaining step and target our loss to our target latents
|
||||
@@ -258,6 +307,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
):
|
||||
loss_target = self.train_config.loss_target
|
||||
is_reg = any(batch.get_is_reg_list())
|
||||
additional_loss = 0.0
|
||||
|
||||
prior_mask_multiplier = None
|
||||
target_mask_multiplier = None
|
||||
@@ -310,24 +360,20 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
if self.train_config.inverted_mask_prior and prior_pred is not None and has_mask:
|
||||
assert not self.train_config.train_turbo
|
||||
with torch.no_grad():
|
||||
# we need to make the noise prediction be a masked blending of noise and prior_pred
|
||||
stretched_mask_multiplier = value_map(
|
||||
mask_multiplier,
|
||||
batch.file_items[0].dataset_config.mask_min_value,
|
||||
1.0,
|
||||
0.0,
|
||||
1.0
|
||||
)
|
||||
prior_mask = batch.mask_tensor.to(self.device_torch, dtype=dtype)
|
||||
# resize to size of noise_pred
|
||||
prior_mask = torch.nn.functional.interpolate(prior_mask, size=(noise_pred.shape[2], noise_pred.shape[3]), mode='bicubic')
|
||||
# stack first channel to match channels of noise_pred
|
||||
prior_mask = torch.cat([prior_mask[:1]] * noise_pred.shape[1], dim=1)
|
||||
|
||||
prior_mask_multiplier = 1.0 - stretched_mask_multiplier
|
||||
|
||||
|
||||
# target_mask_multiplier = mask_multiplier
|
||||
# mask_multiplier = 1.0
|
||||
target = noise
|
||||
# target = (noise * mask_multiplier) + (prior_pred * prior_mask_multiplier)
|
||||
# set masked multiplier to 1.0 so we dont double apply it
|
||||
# mask_multiplier = 1.0
|
||||
prior_mask_multiplier = 1.0 - prior_mask
|
||||
|
||||
# scale so it is a mean of 1
|
||||
prior_mask_multiplier = prior_mask_multiplier / prior_mask_multiplier.mean()
|
||||
if self.sd.is_flow_matching:
|
||||
target = (noise - batch.latents).detach()
|
||||
else:
|
||||
target = noise
|
||||
elif prior_pred is not None and not self.train_config.do_prior_divergence:
|
||||
assert not self.train_config.train_turbo
|
||||
# matching adapter prediction
|
||||
@@ -335,12 +381,62 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
elif self.sd.prediction_type == 'v_prediction':
|
||||
# v-parameterization training
|
||||
target = self.sd.noise_scheduler.get_velocity(batch.tensor, noise, timesteps)
|
||||
|
||||
|
||||
elif hasattr(self.sd, 'get_loss_target'):
|
||||
target = self.sd.get_loss_target(
|
||||
noise=noise,
|
||||
batch=batch,
|
||||
timesteps=timesteps,
|
||||
).detach()
|
||||
|
||||
elif self.sd.is_flow_matching:
|
||||
# forward ODE
|
||||
target = (noise - batch.latents).detach()
|
||||
# reverse ODE
|
||||
# target = (batch.latents - noise).detach()
|
||||
else:
|
||||
target = noise
|
||||
|
||||
|
||||
if self.dfe is not None:
|
||||
if self.dfe.version == 1:
|
||||
# do diffusion feature extraction on target
|
||||
with torch.no_grad():
|
||||
rectified_flow_target = noise.float() - batch.latents.float()
|
||||
target_features = self.dfe(torch.cat([rectified_flow_target, noise.float()], dim=1))
|
||||
|
||||
# do diffusion feature extraction on prediction
|
||||
pred_features = self.dfe(torch.cat([noise_pred.float(), noise.float()], dim=1))
|
||||
additional_loss += torch.nn.functional.mse_loss(pred_features, target_features, reduction="mean") * \
|
||||
self.train_config.diffusion_feature_extractor_weight
|
||||
elif self.dfe.version == 2:
|
||||
# version 2
|
||||
# do diffusion feature extraction on target
|
||||
with torch.no_grad():
|
||||
rectified_flow_target = noise.float() - batch.latents.float()
|
||||
target_feature_list = self.dfe(torch.cat([rectified_flow_target, noise.float()], dim=1))
|
||||
|
||||
# do diffusion feature extraction on prediction
|
||||
pred_feature_list = self.dfe(torch.cat([noise_pred.float(), noise.float()], dim=1))
|
||||
|
||||
dfe_loss = 0.0
|
||||
for i in range(len(target_feature_list)):
|
||||
dfe_loss += torch.nn.functional.mse_loss(pred_feature_list[i], target_feature_list[i], reduction="mean")
|
||||
|
||||
additional_loss += dfe_loss * self.train_config.diffusion_feature_extractor_weight * 100.0
|
||||
elif self.dfe.version == 3:
|
||||
dfe_loss = self.dfe(
|
||||
noise=noise,
|
||||
noise_pred=noise_pred,
|
||||
noisy_latents=noisy_latents,
|
||||
timesteps=timesteps,
|
||||
batch=batch,
|
||||
scheduler=self.sd.noise_scheduler
|
||||
)
|
||||
additional_loss += dfe_loss * self.train_config.diffusion_feature_extractor_weight
|
||||
else:
|
||||
raise ValueError(f"Unknown diffusion feature extractor version {self.dfe.version}")
|
||||
|
||||
|
||||
if target is None:
|
||||
target = noise
|
||||
|
||||
@@ -386,13 +482,18 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
if self.train_config.loss_type == "mae":
|
||||
loss = torch.nn.functional.l1_loss(pred.float(), target.float(), reduction="none")
|
||||
elif self.train_config.loss_type == "wavelet":
|
||||
loss = wavelet_loss(pred, batch.latents, noise)
|
||||
else:
|
||||
loss = torch.nn.functional.mse_loss(pred.float(), target.float(), reduction="none")
|
||||
|
||||
# handle linear timesteps and only adjust the weight of the timesteps
|
||||
if self.sd.is_flow_matching and self.train_config.linear_timesteps:
|
||||
if self.sd.is_flow_matching and (self.train_config.linear_timesteps or self.train_config.linear_timesteps2):
|
||||
# calculate the weights for the timesteps
|
||||
timestep_weight = self.sd.noise_scheduler.get_weights_for_timesteps(timesteps).to(loss.device, dtype=loss.dtype)
|
||||
timestep_weight = self.sd.noise_scheduler.get_weights_for_timesteps(
|
||||
timesteps,
|
||||
v2=self.train_config.linear_timesteps2
|
||||
).to(loss.device, dtype=loss.dtype)
|
||||
timestep_weight = timestep_weight.view(-1, 1, 1, 1).detach()
|
||||
loss = loss * timestep_weight
|
||||
|
||||
@@ -405,7 +506,11 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
mask_multiplier = torch.nn.functional.interpolate(mask_multiplier, size=(pred.shape[2], pred.shape[3]), mode='nearest')
|
||||
|
||||
# multiply by our mask
|
||||
loss = loss * mask_multiplier
|
||||
try:
|
||||
loss = loss * mask_multiplier
|
||||
except:
|
||||
# todo handle mask with video models
|
||||
pass
|
||||
|
||||
prior_loss = None
|
||||
if self.train_config.inverted_mask_prior and prior_pred is not None and prior_mask_multiplier is not None:
|
||||
@@ -417,7 +522,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
prior_loss = prior_loss * prior_mask_multiplier * self.train_config.inverted_mask_prior_multiplier
|
||||
if torch.isnan(prior_loss).any():
|
||||
print("Prior loss is nan")
|
||||
print_acc("Prior loss is nan")
|
||||
prior_loss = None
|
||||
else:
|
||||
prior_loss = prior_loss.mean([1, 2, 3])
|
||||
@@ -426,7 +531,12 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
# loss = loss + prior_loss
|
||||
loss = loss.mean([1, 2, 3])
|
||||
# apply loss multiplier before prior loss
|
||||
loss = loss * loss_multiplier
|
||||
# multiply by our mask
|
||||
try:
|
||||
loss = loss * loss_multiplier
|
||||
except:
|
||||
# todo handle mask with video models
|
||||
pass
|
||||
if prior_loss is not None:
|
||||
loss = loss + prior_loss
|
||||
|
||||
@@ -457,7 +567,20 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
loss = loss + norm_std_loss
|
||||
|
||||
|
||||
return loss
|
||||
return loss + additional_loss
|
||||
|
||||
def get_diff_output_preservation_loss(
|
||||
self,
|
||||
noise_pred: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
noisy_latents: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
mask_multiplier: Union[torch.Tensor, float] = 1.0,
|
||||
prior_pred: Union[torch.Tensor, None] = None,
|
||||
**kwargs
|
||||
):
|
||||
loss_target = self.train_config.loss_target
|
||||
|
||||
def preprocess_batch(self, batch: 'DataLoaderBatchDTO'):
|
||||
return batch
|
||||
@@ -486,7 +609,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
noise=noise,
|
||||
sd=self.sd,
|
||||
unconditional_embeds=unconditional_embeds,
|
||||
scaler=self.scaler,
|
||||
train_config=self.train_config,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
@@ -601,7 +724,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
# loss = self.apply_snr(loss, timesteps)
|
||||
loss = loss.mean()
|
||||
loss.backward()
|
||||
self.accelerator.backward(loss)
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
@@ -756,7 +879,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
# loss = self.apply_snr(loss, timesteps)
|
||||
loss = loss.mean()
|
||||
loss.backward()
|
||||
self.accelerator.backward(loss)
|
||||
|
||||
# detach it so parent class can run backward on no grads without throwing error
|
||||
loss = loss.detach()
|
||||
@@ -794,6 +917,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
was_adapter_active = self.adapter.is_active
|
||||
self.adapter.is_active = False
|
||||
|
||||
if self.train_config.unload_text_encoder and self.adapter is not None and not isinstance(self.adapter, CustomAdapter):
|
||||
raise ValueError("Prior predictions currently do not support unloading text encoder with adapter")
|
||||
# do a prediction here so we can match its output with network multiplier set to 0.0
|
||||
with torch.no_grad():
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
@@ -838,7 +963,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
# self.network.multiplier = 0.0
|
||||
self.sd.unet.eval()
|
||||
|
||||
if self.adapter is not None and isinstance(self.adapter, IPAdapter) and not self.sd.is_flux:
|
||||
if self.adapter is not None and isinstance(self.adapter, IPAdapter) and not self.sd.is_flux and not self.sd.is_lumina2:
|
||||
# we need to remove the image embeds from the prompt except for flux
|
||||
embeds_to_use: PromptEmbeds = embeds_to_use.clone().detach()
|
||||
end_pos = embeds_to_use.text_embeds.shape[1] - self.adapter_config.num_tokens
|
||||
@@ -902,12 +1027,14 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
unconditional_embeddings=unconditional_embeds,
|
||||
timestep=timesteps,
|
||||
guidance_scale=self.train_config.cfg_scale,
|
||||
guidance_embedding_scale=self.train_config.cfg_scale,
|
||||
detach_unconditional=False,
|
||||
rescale_cfg=self.train_config.cfg_rescale,
|
||||
bypass_guidance_embedding=self.train_config.bypass_guidance_embedding,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
def hook_train_loop(self, batch: 'DataLoaderBatchDTO'):
|
||||
def train_single_accumulation(self, batch: DataLoaderBatchDTO):
|
||||
self.timer.start('preprocess_batch')
|
||||
batch = self.preprocess_batch(batch)
|
||||
dtype = get_torch_dtype(self.train_config.dtype)
|
||||
@@ -939,7 +1066,6 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter):
|
||||
# condition the prompt
|
||||
# todo handle more than one adapter image
|
||||
self.adapter.num_control_images = 1
|
||||
conditioned_prompts = self.adapter.condition_prompt(conditioned_prompts)
|
||||
|
||||
network_weight_list = batch.get_network_weight_list()
|
||||
@@ -1016,6 +1142,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
# expand to match latents
|
||||
mask_multiplier = mask_multiplier.expand(-1, noisy_latents.shape[1], -1, -1)
|
||||
mask_multiplier = mask_multiplier.to(self.device_torch, dtype=dtype).detach()
|
||||
# make avg 1.0
|
||||
mask_multiplier = mask_multiplier / mask_multiplier.mean()
|
||||
|
||||
def get_adapter_multiplier():
|
||||
if self.adapter and isinstance(self.adapter, T2IAdapter):
|
||||
@@ -1070,7 +1198,6 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
# set the weights
|
||||
network.multiplier = network_weight_list
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
# activate network if it exits
|
||||
|
||||
@@ -1166,7 +1293,7 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
self.adapter(conditional_clip_embeds)
|
||||
|
||||
# do the custom adapter after the prior prediction
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter) and has_clip_image:
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter) and (has_clip_image or is_reg):
|
||||
quad_count = random.randint(1, 4)
|
||||
self.adapter.train()
|
||||
self.adapter.trigger_pre_te(
|
||||
@@ -1179,7 +1306,30 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
with self.timer('encode_prompt'):
|
||||
unconditional_embeds = None
|
||||
if grad_on_text_encoder:
|
||||
if self.train_config.unload_text_encoder:
|
||||
with torch.set_grad_enabled(False):
|
||||
embeds_to_use = self.cached_blank_embeds.clone().detach().to(
|
||||
self.device_torch, dtype=dtype
|
||||
)
|
||||
if self.cached_trigger_embeds is not None and not is_reg:
|
||||
embeds_to_use = self.cached_trigger_embeds.clone().detach().to(
|
||||
self.device_torch, dtype=dtype
|
||||
)
|
||||
conditional_embeds = concat_prompt_embeds(
|
||||
[embeds_to_use] * noisy_latents.shape[0]
|
||||
)
|
||||
if self.train_config.do_cfg:
|
||||
unconditional_embeds = self.cached_blank_embeds.clone().detach().to(
|
||||
self.device_torch, dtype=dtype
|
||||
)
|
||||
unconditional_embeds = concat_prompt_embeds(
|
||||
[unconditional_embeds] * noisy_latents.shape[0]
|
||||
)
|
||||
|
||||
if isinstance(self.adapter, CustomAdapter):
|
||||
self.adapter.is_unconditional_run = False
|
||||
|
||||
elif grad_on_text_encoder:
|
||||
with torch.set_grad_enabled(True):
|
||||
if isinstance(self.adapter, CustomAdapter):
|
||||
self.adapter.is_unconditional_run = False
|
||||
@@ -1230,17 +1380,39 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
dtype=dtype)
|
||||
if isinstance(self.adapter, CustomAdapter):
|
||||
self.adapter.is_unconditional_run = False
|
||||
|
||||
|
||||
if self.train_config.diff_output_preservation:
|
||||
dop_prompts = [p.replace(self.trigger_word, self.train_config.diff_output_preservation_class) for p in conditioned_prompts]
|
||||
dop_prompts_2 = None
|
||||
if prompt_2 is not None:
|
||||
dop_prompts_2 = [p.replace(self.trigger_word, self.train_config.diff_output_preservation_class) for p in prompt_2]
|
||||
self.diff_output_preservation_embeds = self.sd.encode_prompt(
|
||||
dop_prompts, dop_prompts_2,
|
||||
dropout_prob=self.train_config.prompt_dropout_prob,
|
||||
long_prompts=self.do_long_prompts).to(
|
||||
self.device_torch,
|
||||
dtype=dtype)
|
||||
# detach the embeddings
|
||||
conditional_embeds = conditional_embeds.detach()
|
||||
if self.train_config.do_cfg:
|
||||
unconditional_embeds = unconditional_embeds.detach()
|
||||
|
||||
if self.decorator:
|
||||
conditional_embeds.text_embeds = self.decorator(
|
||||
conditional_embeds.text_embeds
|
||||
)
|
||||
if self.train_config.do_cfg:
|
||||
unconditional_embeds.text_embeds = self.decorator(
|
||||
unconditional_embeds.text_embeds,
|
||||
is_unconditional=True
|
||||
)
|
||||
|
||||
# flush()
|
||||
pred_kwargs = {}
|
||||
|
||||
if has_adapter_img:
|
||||
if (self.adapter and isinstance(self.adapter, T2IAdapter)) or (self.assistant_adapter and isinstance(self.assistant_adapter, T2IAdapter)):
|
||||
if (self.adapter and isinstance(self.adapter, T2IAdapter)) or (
|
||||
self.assistant_adapter and isinstance(self.assistant_adapter, T2IAdapter)):
|
||||
with torch.set_grad_enabled(self.adapter is not None):
|
||||
adapter = self.assistant_adapter if self.assistant_adapter is not None else self.adapter
|
||||
adapter_multiplier = get_adapter_multiplier()
|
||||
@@ -1280,7 +1452,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
if self.train_config.do_cfg:
|
||||
embeds = [
|
||||
load_file(random.choice(batch.clip_image_embeds_unconditional)) for i in range(noisy_latents.shape[0])
|
||||
load_file(random.choice(batch.clip_image_embeds_unconditional)) for i in
|
||||
range(noisy_latents.shape[0])
|
||||
]
|
||||
unconditional_clip_embeds = self.adapter.parse_clip_image_embeds_from_cache(
|
||||
embeds,
|
||||
@@ -1341,8 +1514,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
quad_count=quad_count
|
||||
)
|
||||
else:
|
||||
print("No Clip Image")
|
||||
print([file_item.path for file_item in batch.file_items])
|
||||
print_acc("No Clip Image")
|
||||
print_acc([file_item.path for file_item in batch.file_items])
|
||||
raise ValueError("Could not find clip image")
|
||||
|
||||
if not self.adapter_config.train_image_encoder:
|
||||
@@ -1406,9 +1579,14 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
if ((
|
||||
has_adapter_img and self.assistant_adapter and match_adapter_assist) or self.do_prior_prediction or do_guidance_prior or do_reg_prior or do_inverted_masked_prior or self.train_config.correct_pred_norm):
|
||||
with self.timer('prior predict'):
|
||||
prior_embeds_to_use = conditional_embeds
|
||||
# use diff_output_preservation embeds if doing dfe
|
||||
if self.train_config.diff_output_preservation:
|
||||
prior_embeds_to_use = self.diff_output_preservation_embeds.expand_to_batch(noisy_latents.shape[0])
|
||||
|
||||
prior_pred = self.get_prior_prediction(
|
||||
noisy_latents=noisy_latents,
|
||||
conditional_embeds=conditional_embeds,
|
||||
conditional_embeds=prior_embeds_to_use,
|
||||
match_adapter_assist=match_adapter_assist,
|
||||
network_weight_list=network_weight_list,
|
||||
timesteps=timesteps,
|
||||
@@ -1421,9 +1599,8 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
if prior_pred is not None:
|
||||
prior_pred = prior_pred.detach()
|
||||
|
||||
|
||||
# do the custom adapter after the prior prediction
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter) and has_clip_image:
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter) and (has_clip_image or self.adapter_config.type in ['llm_adapter', 'text_encoder']):
|
||||
quad_count = random.randint(1, 4)
|
||||
self.adapter.train()
|
||||
conditional_embeds = self.adapter.condition_encoded_embeds(
|
||||
@@ -1447,10 +1624,12 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
self.adapter.add_extra_values(batch.extra_values.detach())
|
||||
|
||||
if self.train_config.do_cfg:
|
||||
self.adapter.add_extra_values(torch.zeros_like(batch.extra_values.detach()), is_unconditional=True)
|
||||
self.adapter.add_extra_values(torch.zeros_like(batch.extra_values.detach()),
|
||||
is_unconditional=True)
|
||||
|
||||
if has_adapter_img:
|
||||
if (self.adapter and isinstance(self.adapter, ControlNetModel)) or (self.assistant_adapter and isinstance(self.assistant_adapter, ControlNetModel)):
|
||||
if (self.adapter and isinstance(self.adapter, ControlNetModel)) or (
|
||||
self.assistant_adapter and isinstance(self.assistant_adapter, ControlNetModel)):
|
||||
if self.train_config.do_cfg:
|
||||
raise ValueError("ControlNetModel is not supported with CFG")
|
||||
with torch.set_grad_enabled(self.adapter is not None):
|
||||
@@ -1475,7 +1654,6 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
pred_kwargs['down_block_additional_residuals'] = down_block_res_samples
|
||||
pred_kwargs['mid_block_additional_residual'] = mid_block_res_sample
|
||||
|
||||
|
||||
self.before_unet_predict()
|
||||
# do a prior pred if we have an unconditional image, we will swap out the giadance later
|
||||
if batch.unconditional_latents is not None or self.do_guided_loss:
|
||||
@@ -1495,9 +1673,12 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
)
|
||||
|
||||
else:
|
||||
if unconditional_embeds is not None:
|
||||
unconditional_embeds = unconditional_embeds.to(self.device_torch, dtype=dtype).detach()
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter):
|
||||
with self.timer('condition_noisy_latents'):
|
||||
noisy_latents = self.adapter.condition_noisy_latents(noisy_latents, batch)
|
||||
with self.timer('predict_unet'):
|
||||
if unconditional_embeds is not None:
|
||||
unconditional_embeds = unconditional_embeds.to(self.device_torch, dtype=dtype).detach()
|
||||
noise_pred = self.predict_noise(
|
||||
noisy_latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
timesteps=timesteps,
|
||||
@@ -1509,6 +1690,12 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
|
||||
with self.timer('calculate_loss'):
|
||||
noise = noise.to(self.device_torch, dtype=dtype).detach()
|
||||
prior_to_calculate_loss = prior_pred
|
||||
# if we are doing diff_output_preservation and not noing inverted masked prior
|
||||
# then we need to send none here so it will not target the prior
|
||||
if self.train_config.diff_output_preservation and not do_inverted_masked_prior:
|
||||
prior_to_calculate_loss = None
|
||||
|
||||
loss = self.calculate_loss(
|
||||
noise_pred=noise_pred,
|
||||
noise=noise,
|
||||
@@ -1516,14 +1703,35 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
timesteps=timesteps,
|
||||
batch=batch,
|
||||
mask_multiplier=mask_multiplier,
|
||||
prior_pred=prior_pred,
|
||||
prior_pred=prior_to_calculate_loss,
|
||||
)
|
||||
|
||||
if self.train_config.diff_output_preservation:
|
||||
# send the loss backwards otherwise checkpointing will fail
|
||||
self.accelerator.backward(loss)
|
||||
normal_loss = loss.detach() # dont send backward again
|
||||
|
||||
dop_embeds = self.diff_output_preservation_embeds.expand_to_batch(noisy_latents.shape[0])
|
||||
dop_pred = self.predict_noise(
|
||||
noisy_latents=noisy_latents.to(self.device_torch, dtype=dtype),
|
||||
timesteps=timesteps,
|
||||
conditional_embeds=dop_embeds.to(self.device_torch, dtype=dtype),
|
||||
unconditional_embeds=unconditional_embeds,
|
||||
**pred_kwargs
|
||||
)
|
||||
dop_loss = torch.nn.functional.mse_loss(dop_pred, prior_pred) * self.train_config.diff_output_preservation_multiplier
|
||||
self.accelerator.backward(dop_loss)
|
||||
|
||||
loss = normal_loss + dop_loss
|
||||
loss = loss.clone().detach()
|
||||
# require grad again so the backward wont fail
|
||||
loss.requires_grad_(True)
|
||||
|
||||
# check if nan
|
||||
if torch.isnan(loss):
|
||||
print("loss is nan")
|
||||
print_acc("loss is nan")
|
||||
loss = torch.zeros_like(loss).requires_grad_(True)
|
||||
|
||||
|
||||
with self.timer('backward'):
|
||||
# todo we have multiplier seperated. works for now as res are not in same batch, but need to change
|
||||
loss = loss * loss_multiplier.mean()
|
||||
@@ -1536,32 +1744,43 @@ class SDTrainer(BaseSDTrainProcess):
|
||||
# if self.is_bfloat:
|
||||
# loss.backward()
|
||||
# else:
|
||||
if not self.do_grad_scale:
|
||||
loss.backward()
|
||||
else:
|
||||
self.scaler.scale(loss).backward()
|
||||
self.accelerator.backward(loss)
|
||||
|
||||
return loss.detach()
|
||||
# flush()
|
||||
|
||||
def hook_train_loop(self, batch: Union[DataLoaderBatchDTO, List[DataLoaderBatchDTO]]):
|
||||
if isinstance(batch, list):
|
||||
batch_list = batch
|
||||
else:
|
||||
batch_list = [batch]
|
||||
total_loss = None
|
||||
self.optimizer.zero_grad()
|
||||
for batch in batch_list:
|
||||
loss = self.train_single_accumulation(batch)
|
||||
if total_loss is None:
|
||||
total_loss = loss
|
||||
else:
|
||||
total_loss += loss
|
||||
if len(batch_list) > 1 and self.model_config.low_vram:
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
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)
|
||||
self.accelerator.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)
|
||||
self.accelerator.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.step()
|
||||
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
if self.adapter and isinstance(self.adapter, CustomAdapter):
|
||||
self.adapter.post_weight_update()
|
||||
if self.ema is not None:
|
||||
with self.timer('ema_update'):
|
||||
self.ema.update()
|
||||
|
||||
234
extensions_built_in/sd_trainer/UITrainer.py
Normal file
234
extensions_built_in/sd_trainer/UITrainer.py
Normal file
@@ -0,0 +1,234 @@
|
||||
from collections import OrderedDict
|
||||
import os
|
||||
import sqlite3
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
from extensions_built_in.sd_trainer.SDTrainer import SDTrainer
|
||||
from typing import Literal, Optional
|
||||
|
||||
|
||||
AITK_Status = Literal["running", "stopped", "error", "completed"]
|
||||
|
||||
|
||||
class UITrainer(SDTrainer):
|
||||
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
|
||||
super(UITrainer, self).__init__(process_id, job, config, **kwargs)
|
||||
self.sqlite_db_path = self.config.get("sqlite_db_path", "./aitk_db.db")
|
||||
if not os.path.exists(self.sqlite_db_path):
|
||||
raise Exception(
|
||||
f"SQLite database not found at {self.sqlite_db_path}")
|
||||
print(f"Using SQLite database at {self.sqlite_db_path}")
|
||||
self.job_id = os.environ.get("AITK_JOB_ID", None)
|
||||
self.job_id = self.job_id.strip() if self.job_id is not None else None
|
||||
print(f"Job ID: \"{self.job_id}\"")
|
||||
if self.job_id is None:
|
||||
raise Exception("AITK_JOB_ID not set")
|
||||
self.is_stopping = False
|
||||
# Create a thread pool for database operations
|
||||
self.thread_pool = concurrent.futures.ThreadPoolExecutor(max_workers=1)
|
||||
# Track all async tasks
|
||||
self._async_tasks = []
|
||||
# Initialize the status
|
||||
self._run_async_operation(self._update_status("running", "Starting"))
|
||||
|
||||
def _run_async_operation(self, coro):
|
||||
"""Helper method to run an async coroutine and track the task."""
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
except RuntimeError:
|
||||
# No event loop exists, create a new one
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
|
||||
# Create a task and track it
|
||||
if loop.is_running():
|
||||
task = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||
self._async_tasks.append(asyncio.wrap_future(task))
|
||||
else:
|
||||
task = loop.create_task(coro)
|
||||
self._async_tasks.append(task)
|
||||
loop.run_until_complete(task)
|
||||
|
||||
async def _execute_db_operation(self, operation_func):
|
||||
"""Execute a database operation in a separate thread to avoid blocking."""
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(self.thread_pool, operation_func)
|
||||
|
||||
def _db_connect(self):
|
||||
"""Create a new connection for each operation to avoid locking."""
|
||||
conn = sqlite3.connect(self.sqlite_db_path, timeout=10.0)
|
||||
conn.isolation_level = None # Enable autocommit mode
|
||||
return conn
|
||||
|
||||
def should_stop(self):
|
||||
def _check_stop():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute(
|
||||
"SELECT stop FROM Job WHERE id = ?", (self.job_id,))
|
||||
stop = cursor.fetchone()
|
||||
return False if stop is None else stop[0] == 1
|
||||
|
||||
return _check_stop()
|
||||
|
||||
def maybe_stop(self):
|
||||
if self.should_stop():
|
||||
self._run_async_operation(
|
||||
self._update_status("stopped", "Job stopped"))
|
||||
self.is_stopping = True
|
||||
raise Exception("Job stopped")
|
||||
|
||||
async def _update_key(self, key, value):
|
||||
if not self.accelerator.is_main_process:
|
||||
return
|
||||
|
||||
def _do_update():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
# Convert the value to string if it's not already
|
||||
if isinstance(value, str):
|
||||
value_to_insert = value
|
||||
else:
|
||||
value_to_insert = str(value)
|
||||
|
||||
# Use parameterized query for both the column name and value
|
||||
update_query = f"UPDATE Job SET {key} = ? WHERE id = ?"
|
||||
cursor.execute(
|
||||
update_query, (value_to_insert, self.job_id))
|
||||
finally:
|
||||
cursor.execute("COMMIT")
|
||||
|
||||
await self._execute_db_operation(_do_update)
|
||||
|
||||
def update_step(self):
|
||||
"""Non-blocking update of the step count."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_key("step", self.step_num))
|
||||
|
||||
def update_db_key(self, key, value):
|
||||
"""Non-blocking update a key in the database."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_key(key, value))
|
||||
|
||||
async def _update_status(self, status: AITK_Status, info: Optional[str] = None):
|
||||
if not self.accelerator.is_main_process:
|
||||
return
|
||||
|
||||
def _do_update():
|
||||
with self._db_connect() as conn:
|
||||
cursor = conn.cursor()
|
||||
cursor.execute("BEGIN IMMEDIATE")
|
||||
try:
|
||||
if info is not None:
|
||||
cursor.execute(
|
||||
"UPDATE Job SET status = ?, info = ? WHERE id = ?",
|
||||
(status, info, self.job_id)
|
||||
)
|
||||
else:
|
||||
cursor.execute(
|
||||
"UPDATE Job SET status = ? WHERE id = ?",
|
||||
(status, self.job_id)
|
||||
)
|
||||
finally:
|
||||
cursor.execute("COMMIT")
|
||||
|
||||
await self._execute_db_operation(_do_update)
|
||||
|
||||
def update_status(self, status: AITK_Status, info: Optional[str] = None):
|
||||
"""Non-blocking update of status."""
|
||||
if self.accelerator.is_main_process:
|
||||
self._run_async_operation(self._update_status(status, info))
|
||||
|
||||
async def wait_for_all_async(self):
|
||||
"""Wait for all tracked async operations to complete."""
|
||||
if not self._async_tasks:
|
||||
return
|
||||
|
||||
try:
|
||||
await asyncio.gather(*self._async_tasks)
|
||||
except Exception as e:
|
||||
pass
|
||||
finally:
|
||||
# Clear the task list after completion
|
||||
self._async_tasks.clear()
|
||||
|
||||
def on_error(self, e: Exception):
|
||||
super(UITrainer, self).on_error(e)
|
||||
if self.accelerator.is_main_process and not self.is_stopping:
|
||||
self.update_status("error", str(e))
|
||||
self.update_db_key("step", self.last_save_step)
|
||||
asyncio.run(self.wait_for_all_async())
|
||||
self.thread_pool.shutdown(wait=True)
|
||||
|
||||
def handle_timing_print_hook(self, timing_dict):
|
||||
if "train_loop" not in timing_dict:
|
||||
print("train_loop not found in timing_dict", timing_dict)
|
||||
return
|
||||
seconds_per_iter = timing_dict["train_loop"]
|
||||
# determine iter/sec or sec/iter
|
||||
if seconds_per_iter < 1:
|
||||
iters_per_sec = 1 / seconds_per_iter
|
||||
self.update_db_key("speed_string", f"{iters_per_sec:.2f} iter/sec")
|
||||
else:
|
||||
self.update_db_key(
|
||||
"speed_string", f"{seconds_per_iter:.2f} sec/iter")
|
||||
|
||||
def done_hook(self):
|
||||
super(UITrainer, self).done_hook()
|
||||
self.update_status("completed", "Training completed")
|
||||
# Wait for all async operations to finish before shutting down
|
||||
asyncio.run(self.wait_for_all_async())
|
||||
self.thread_pool.shutdown(wait=True)
|
||||
|
||||
def end_step_hook(self):
|
||||
super(UITrainer, self).end_step_hook()
|
||||
self.update_step()
|
||||
self.maybe_stop()
|
||||
|
||||
def hook_before_model_load(self):
|
||||
super().hook_before_model_load()
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Loading model")
|
||||
|
||||
def before_dataset_load(self):
|
||||
super().before_dataset_load()
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Loading dataset")
|
||||
|
||||
def hook_before_train_loop(self):
|
||||
super().hook_before_train_loop()
|
||||
self.maybe_stop()
|
||||
self.update_step()
|
||||
self.update_status("running", "Training")
|
||||
self.timer.add_after_print_hook(self.handle_timing_print_hook)
|
||||
|
||||
def status_update_hook_func(self, string):
|
||||
self.update_status("running", string)
|
||||
|
||||
def hook_after_sd_init_before_load(self):
|
||||
super().hook_after_sd_init_before_load()
|
||||
self.maybe_stop()
|
||||
self.sd.add_status_update_hook(self.status_update_hook_func)
|
||||
|
||||
def sample_step_hook(self, img_num, total_imgs):
|
||||
super().sample_step_hook(img_num, total_imgs)
|
||||
self.maybe_stop()
|
||||
self.update_status(
|
||||
"running", f"Generating images - {img_num + 1}/{total_imgs}")
|
||||
|
||||
def sample(self, step=None, is_first=False):
|
||||
self.maybe_stop()
|
||||
total_imgs = len(self.sample_config.prompts)
|
||||
self.update_status("running", f"Generating images - 0/{total_imgs}")
|
||||
super().sample(step, is_first)
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Training")
|
||||
|
||||
def save(self, step=None):
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Saving model")
|
||||
super().save(step)
|
||||
self.maybe_stop()
|
||||
self.update_status("running", "Training")
|
||||
@@ -18,6 +18,22 @@ class SDTrainerExtension(Extension):
|
||||
from .SDTrainer import SDTrainer
|
||||
return SDTrainer
|
||||
|
||||
# This is for generic training (LoRA, Dreambooth, FineTuning)
|
||||
class UITrainerExtension(Extension):
|
||||
# uid must be unique, it is how the extension is identified
|
||||
uid = "ui_trainer"
|
||||
|
||||
# name is the name of the extension for printing
|
||||
name = "UI Trainer"
|
||||
|
||||
# This is where your process class is loaded
|
||||
# keep your imports in here so they don't slow down the rest of the program
|
||||
@classmethod
|
||||
def get_process(cls):
|
||||
# import your process class here so it is only loaded when needed and return it
|
||||
from .UITrainer import UITrainer
|
||||
return UITrainer
|
||||
|
||||
|
||||
# for backwards compatability
|
||||
class TextualInversionTrainer(SDTrainerExtension):
|
||||
@@ -26,5 +42,5 @@ class TextualInversionTrainer(SDTrainerExtension):
|
||||
|
||||
AI_TOOLKIT_EXTENSIONS = [
|
||||
# you can put a list of extensions here
|
||||
SDTrainerExtension, TextualInversionTrainer
|
||||
SDTrainerExtension, TextualInversionTrainer, UITrainerExtension
|
||||
]
|
||||
|
||||
414
flux_train_ui.py
Normal file
414
flux_train_ui.py
Normal file
@@ -0,0 +1,414 @@
|
||||
import os
|
||||
from huggingface_hub import whoami
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
|
||||
# Add the current working directory to the Python path
|
||||
sys.path.insert(0, os.getcwd())
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
import torch
|
||||
import uuid
|
||||
import os
|
||||
import shutil
|
||||
import json
|
||||
import yaml
|
||||
from slugify import slugify
|
||||
from transformers import AutoProcessor, AutoModelForCausalLM
|
||||
|
||||
sys.path.insert(0, "ai-toolkit")
|
||||
from toolkit.job import get_job
|
||||
|
||||
MAX_IMAGES = 150
|
||||
|
||||
def load_captioning(uploaded_files, concept_sentence):
|
||||
uploaded_images = [file for file in uploaded_files if not file.endswith('.txt')]
|
||||
txt_files = [file for file in uploaded_files if file.endswith('.txt')]
|
||||
txt_files_dict = {os.path.splitext(os.path.basename(txt_file))[0]: txt_file for txt_file in txt_files}
|
||||
updates = []
|
||||
if len(uploaded_images) <= 1:
|
||||
raise gr.Error(
|
||||
"Please upload at least 2 images to train your model (the ideal number with default settings is between 4-30)"
|
||||
)
|
||||
elif len(uploaded_images) > MAX_IMAGES:
|
||||
raise gr.Error(f"For now, only {MAX_IMAGES} or less images are allowed for training")
|
||||
# Update for the captioning_area
|
||||
# for _ in range(3):
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update visibility and image for each captioning row and image
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
# Determine if the current row and image should be visible
|
||||
visible = i <= len(uploaded_images)
|
||||
|
||||
# Update visibility of the captioning row
|
||||
updates.append(gr.update(visible=visible))
|
||||
|
||||
# Update for image component - display image if available, otherwise hide
|
||||
image_value = uploaded_images[i - 1] if visible else None
|
||||
updates.append(gr.update(value=image_value, visible=visible))
|
||||
|
||||
corresponding_caption = False
|
||||
if(image_value):
|
||||
base_name = os.path.splitext(os.path.basename(image_value))[0]
|
||||
print(base_name)
|
||||
print(image_value)
|
||||
if base_name in txt_files_dict:
|
||||
print("entrou")
|
||||
with open(txt_files_dict[base_name], 'r') as file:
|
||||
corresponding_caption = file.read()
|
||||
|
||||
# Update value of captioning area
|
||||
text_value = corresponding_caption if visible and corresponding_caption else "[trigger]" if visible and concept_sentence else None
|
||||
updates.append(gr.update(value=text_value, visible=visible))
|
||||
|
||||
# Update for the sample caption area
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update prompt samples
|
||||
updates.append(gr.update(placeholder=f'A portrait of person in a bustling cafe {concept_sentence}', value=f'A person in a bustling cafe {concept_sentence}'))
|
||||
updates.append(gr.update(placeholder=f"A mountainous landscape in the style of {concept_sentence}"))
|
||||
updates.append(gr.update(placeholder=f"A {concept_sentence} in a mall"))
|
||||
updates.append(gr.update(visible=True))
|
||||
return updates
|
||||
|
||||
def hide_captioning():
|
||||
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
|
||||
|
||||
def create_dataset(*inputs):
|
||||
print("Creating dataset")
|
||||
images = inputs[0]
|
||||
destination_folder = str(f"datasets/{uuid.uuid4()}")
|
||||
if not os.path.exists(destination_folder):
|
||||
os.makedirs(destination_folder)
|
||||
|
||||
jsonl_file_path = os.path.join(destination_folder, "metadata.jsonl")
|
||||
with open(jsonl_file_path, "a") as jsonl_file:
|
||||
for index, image in enumerate(images):
|
||||
new_image_path = shutil.copy(image, destination_folder)
|
||||
|
||||
original_caption = inputs[index + 1]
|
||||
file_name = os.path.basename(new_image_path)
|
||||
|
||||
data = {"file_name": file_name, "prompt": original_caption}
|
||||
|
||||
jsonl_file.write(json.dumps(data) + "\n")
|
||||
|
||||
return destination_folder
|
||||
|
||||
|
||||
def run_captioning(images, concept_sentence, *captions):
|
||||
#Load internally to not consume resources for training
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
torch_dtype = torch.float16
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
"multimodalart/Florence-2-large-no-flash-attn", torch_dtype=torch_dtype, trust_remote_code=True
|
||||
).to(device)
|
||||
processor = AutoProcessor.from_pretrained("multimodalart/Florence-2-large-no-flash-attn", trust_remote_code=True)
|
||||
|
||||
captions = list(captions)
|
||||
for i, image_path in enumerate(images):
|
||||
print(captions[i])
|
||||
if isinstance(image_path, str): # If image is a file path
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
|
||||
prompt = "<DETAILED_CAPTION>"
|
||||
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
|
||||
|
||||
generated_ids = model.generate(
|
||||
input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=1024, num_beams=3
|
||||
)
|
||||
|
||||
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
|
||||
parsed_answer = processor.post_process_generation(
|
||||
generated_text, task=prompt, image_size=(image.width, image.height)
|
||||
)
|
||||
caption_text = parsed_answer["<DETAILED_CAPTION>"].replace("The image shows ", "")
|
||||
if concept_sentence:
|
||||
caption_text = f"{caption_text} [trigger]"
|
||||
captions[i] = caption_text
|
||||
|
||||
yield captions
|
||||
model.to("cpu")
|
||||
del model
|
||||
del processor
|
||||
|
||||
def recursive_update(d, u):
|
||||
for k, v in u.items():
|
||||
if isinstance(v, dict) and v:
|
||||
d[k] = recursive_update(d.get(k, {}), v)
|
||||
else:
|
||||
d[k] = v
|
||||
return d
|
||||
|
||||
def start_training(
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options,
|
||||
):
|
||||
push_to_hub = True
|
||||
if not lora_name:
|
||||
raise gr.Error("You forgot to insert your LoRA name! This name has to be unique.")
|
||||
try:
|
||||
if whoami()["auth"]["accessToken"]["role"] == "write" or "repo.write" in whoami()["auth"]["accessToken"]["fineGrained"]["scoped"][0]["permissions"]:
|
||||
gr.Info(f"Starting training locally {whoami()['name']}. Your LoRA will be available locally and in Hugging Face after it finishes.")
|
||||
else:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
except:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
|
||||
print("Started training")
|
||||
slugged_lora_name = slugify(lora_name)
|
||||
|
||||
# Load the default config
|
||||
with open("config/examples/train_lora_flux_24gb.yaml", "r") as f:
|
||||
config = yaml.safe_load(f)
|
||||
|
||||
# Update the config with user inputs
|
||||
config["config"]["name"] = slugged_lora_name
|
||||
config["config"]["process"][0]["model"]["low_vram"] = low_vram
|
||||
config["config"]["process"][0]["train"]["skip_first_sample"] = True
|
||||
config["config"]["process"][0]["train"]["steps"] = int(steps)
|
||||
config["config"]["process"][0]["train"]["lr"] = float(lr)
|
||||
config["config"]["process"][0]["network"]["linear"] = int(rank)
|
||||
config["config"]["process"][0]["network"]["linear_alpha"] = int(rank)
|
||||
config["config"]["process"][0]["datasets"][0]["folder_path"] = dataset_folder
|
||||
config["config"]["process"][0]["save"]["push_to_hub"] = push_to_hub
|
||||
if(push_to_hub):
|
||||
try:
|
||||
username = whoami()["name"]
|
||||
except:
|
||||
raise gr.Error("Error trying to retrieve your username. Are you sure you are logged in with Hugging Face?")
|
||||
config["config"]["process"][0]["save"]["hf_repo_id"] = f"{username}/{slugged_lora_name}"
|
||||
config["config"]["process"][0]["save"]["hf_private"] = True
|
||||
if concept_sentence:
|
||||
config["config"]["process"][0]["trigger_word"] = concept_sentence
|
||||
|
||||
if sample_1 or sample_2 or sample_3:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = False
|
||||
config["config"]["process"][0]["sample"]["sample_every"] = steps
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 28
|
||||
config["config"]["process"][0]["sample"]["prompts"] = []
|
||||
if sample_1:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_1)
|
||||
if sample_2:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_2)
|
||||
if sample_3:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_3)
|
||||
else:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = True
|
||||
if(model_to_train == "schnell"):
|
||||
config["config"]["process"][0]["model"]["name_or_path"] = "black-forest-labs/FLUX.1-schnell"
|
||||
config["config"]["process"][0]["model"]["assistant_lora_path"] = "ostris/FLUX.1-schnell-training-adapter"
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 4
|
||||
if(use_more_advanced_options):
|
||||
more_advanced_options_dict = yaml.safe_load(more_advanced_options)
|
||||
config["config"]["process"][0] = recursive_update(config["config"]["process"][0], more_advanced_options_dict)
|
||||
print(config)
|
||||
|
||||
# Save the updated config
|
||||
# generate a random name for the config
|
||||
random_config_name = str(uuid.uuid4())
|
||||
os.makedirs("tmp", exist_ok=True)
|
||||
config_path = f"tmp/{random_config_name}-{slugged_lora_name}.yaml"
|
||||
with open(config_path, "w") as f:
|
||||
yaml.dump(config, f)
|
||||
|
||||
# run the job locally
|
||||
job = get_job(config_path)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
|
||||
return f"Training completed successfully. Model saved as {slugged_lora_name}"
|
||||
|
||||
config_yaml = '''
|
||||
device: cuda:0
|
||||
model:
|
||||
is_flux: true
|
||||
quantize: true
|
||||
network:
|
||||
linear: 16 #it will overcome the 'rank' parameter
|
||||
linear_alpha: 16 #you can have an alpha different than the ranking if you'd like
|
||||
type: lora
|
||||
sample:
|
||||
guidance_scale: 3.5
|
||||
height: 1024
|
||||
neg: '' #doesn't work for FLUX
|
||||
sample_every: 1000
|
||||
sample_steps: 28
|
||||
sampler: flowmatch
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
width: 1024
|
||||
save:
|
||||
dtype: float16
|
||||
hf_private: true
|
||||
max_step_saves_to_keep: 4
|
||||
push_to_hub: true
|
||||
save_every: 10000
|
||||
train:
|
||||
batch_size: 1
|
||||
dtype: bf16
|
||||
ema_config:
|
||||
ema_decay: 0.99
|
||||
use_ema: true
|
||||
gradient_accumulation_steps: 1
|
||||
gradient_checkpointing: true
|
||||
noise_scheduler: flowmatch
|
||||
optimizer: adamw8bit #options: prodigy, dadaptation, adamw, adamw8bit, lion, lion8bit
|
||||
train_text_encoder: false #probably doesn't work for flux
|
||||
train_unet: true
|
||||
'''
|
||||
|
||||
theme = gr.themes.Monochrome(
|
||||
text_size=gr.themes.Size(lg="18px", md="15px", sm="13px", xl="22px", xs="12px", xxl="24px", xxs="9px"),
|
||||
font=[gr.themes.GoogleFont("Source Sans Pro"), "ui-sans-serif", "system-ui", "sans-serif"],
|
||||
)
|
||||
css = """
|
||||
h1{font-size: 2em}
|
||||
h3{margin-top: 0}
|
||||
#component-1{text-align:center}
|
||||
.main_ui_logged_out{opacity: 0.3; pointer-events: none}
|
||||
.tabitem{border: 0px}
|
||||
.group_padding{padding: .55em}
|
||||
"""
|
||||
with gr.Blocks(theme=theme, css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""# LoRA Ease for FLUX 🧞♂️
|
||||
### Train a high quality FLUX LoRA in a breeze ༄ using [Ostris' AI Toolkit](https://github.com/ostris/ai-toolkit)"""
|
||||
)
|
||||
with gr.Column() as main_ui:
|
||||
with gr.Row():
|
||||
lora_name = gr.Textbox(
|
||||
label="The name of your LoRA",
|
||||
info="This has to be a unique name",
|
||||
placeholder="e.g.: Persian Miniature Painting style, Cat Toy",
|
||||
)
|
||||
concept_sentence = gr.Textbox(
|
||||
label="Trigger word/sentence",
|
||||
info="Trigger word or sentence to be used",
|
||||
placeholder="uncommon word like p3rs0n or trtcrd, or sentence like 'in the style of CNSTLL'",
|
||||
interactive=True,
|
||||
)
|
||||
with gr.Group(visible=True) as image_upload:
|
||||
with gr.Row():
|
||||
images = gr.File(
|
||||
file_types=["image", ".txt"],
|
||||
label="Upload your images",
|
||||
file_count="multiple",
|
||||
interactive=True,
|
||||
visible=True,
|
||||
scale=1,
|
||||
)
|
||||
with gr.Column(scale=3, visible=False) as captioning_area:
|
||||
with gr.Column():
|
||||
gr.Markdown(
|
||||
"""# Custom captioning
|
||||
<p style="margin-top:0">You can optionally add a custom caption for each image (or use an AI model for this). [trigger] will represent your concept sentence/trigger word.</p>
|
||||
""", elem_classes="group_padding")
|
||||
do_captioning = gr.Button("Add AI captions with Florence-2")
|
||||
output_components = [captioning_area]
|
||||
caption_list = []
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
locals()[f"captioning_row_{i}"] = gr.Row(visible=False)
|
||||
with locals()[f"captioning_row_{i}"]:
|
||||
locals()[f"image_{i}"] = gr.Image(
|
||||
type="filepath",
|
||||
width=111,
|
||||
height=111,
|
||||
min_width=111,
|
||||
interactive=False,
|
||||
scale=2,
|
||||
show_label=False,
|
||||
show_share_button=False,
|
||||
show_download_button=False,
|
||||
)
|
||||
locals()[f"caption_{i}"] = gr.Textbox(
|
||||
label=f"Caption {i}", scale=15, interactive=True
|
||||
)
|
||||
|
||||
output_components.append(locals()[f"captioning_row_{i}"])
|
||||
output_components.append(locals()[f"image_{i}"])
|
||||
output_components.append(locals()[f"caption_{i}"])
|
||||
caption_list.append(locals()[f"caption_{i}"])
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
steps = gr.Number(label="Steps", value=1000, minimum=1, maximum=10000, step=1)
|
||||
lr = gr.Number(label="Learning Rate", value=4e-4, minimum=1e-6, maximum=1e-3, step=1e-6)
|
||||
rank = gr.Number(label="LoRA Rank", value=16, minimum=4, maximum=128, step=4)
|
||||
model_to_train = gr.Radio(["dev", "schnell"], value="dev", label="Model to train")
|
||||
low_vram = gr.Checkbox(label="Low VRAM", value=True)
|
||||
with gr.Accordion("Even more advanced options", open=False):
|
||||
use_more_advanced_options = gr.Checkbox(label="Use more advanced options", value=False)
|
||||
more_advanced_options = gr.Code(config_yaml, language="yaml")
|
||||
|
||||
with gr.Accordion("Sample prompts (optional)", visible=False) as sample:
|
||||
gr.Markdown(
|
||||
"Include sample prompts to test out your trained model. Don't forget to include your trigger word/sentence (optional)"
|
||||
)
|
||||
sample_1 = gr.Textbox(label="Test prompt 1")
|
||||
sample_2 = gr.Textbox(label="Test prompt 2")
|
||||
sample_3 = gr.Textbox(label="Test prompt 3")
|
||||
|
||||
output_components.append(sample)
|
||||
output_components.append(sample_1)
|
||||
output_components.append(sample_2)
|
||||
output_components.append(sample_3)
|
||||
start = gr.Button("Start training", visible=False)
|
||||
output_components.append(start)
|
||||
progress_area = gr.Markdown("")
|
||||
|
||||
dataset_folder = gr.State()
|
||||
|
||||
images.upload(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.delete(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.clear(
|
||||
hide_captioning,
|
||||
outputs=[captioning_area, sample, start]
|
||||
)
|
||||
|
||||
start.click(fn=create_dataset, inputs=[images] + caption_list, outputs=dataset_folder).then(
|
||||
fn=start_training,
|
||||
inputs=[
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options
|
||||
],
|
||||
outputs=progress_area,
|
||||
)
|
||||
|
||||
do_captioning.click(fn=run_captioning, inputs=[images, concept_sentence] + caption_list, outputs=caption_list)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.launch(share=True, show_error=True)
|
||||
@@ -24,6 +24,9 @@ class BaseProcess(object):
|
||||
self.performance_log_every = self.get_conf('performance_log_every', 0)
|
||||
|
||||
print(json.dumps(self.config, indent=4))
|
||||
|
||||
def on_error(self, e: Exception):
|
||||
pass
|
||||
|
||||
def get_conf(self, key, default=None, required=False, as_type=None):
|
||||
# split key by '.' and recursively get the value
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -34,6 +34,7 @@ class GenerateConfig:
|
||||
self.compile = kwargs.get('compile', False)
|
||||
self.ext = kwargs.get('ext', 'png')
|
||||
self.prompt_file = kwargs.get('prompt_file', False)
|
||||
self.num_repeats = kwargs.get('num_repeats', 1)
|
||||
self.prompts_in_file = self.prompts
|
||||
if self.prompts is None:
|
||||
raise ValueError("Prompts must be set")
|
||||
@@ -110,30 +111,31 @@ class GenerateProcess(BaseProcess):
|
||||
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)
|
||||
for _ in range(self.generate_config.num_repeats):
|
||||
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)
|
||||
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
|
||||
))
|
||||
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)
|
||||
|
||||
|
||||
@@ -275,6 +275,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
return adapter_tensors
|
||||
|
||||
def hook_train_loop(self, batch: Union['DataLoaderBatchDTO', None]):
|
||||
if isinstance(batch, list):
|
||||
batch = batch[0]
|
||||
# set to eval mode
|
||||
self.sd.set_device_state(self.eval_slider_device_state)
|
||||
with torch.no_grad():
|
||||
@@ -364,10 +366,32 @@ class TrainSliderProcess(BaseSDTrainProcess):
|
||||
denoised_latents = torch.cat([noisy_latents] * self.prompt_chunk_size, dim=0)
|
||||
current_timestep = timesteps
|
||||
else:
|
||||
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
if self.train_config.noise_scheduler == 'flowmatch':
|
||||
linear_timesteps = any([
|
||||
self.train_config.linear_timesteps,
|
||||
self.train_config.linear_timesteps2,
|
||||
self.train_config.timestep_type == 'linear',
|
||||
])
|
||||
|
||||
timestep_type = 'linear' if linear_timesteps else None
|
||||
if timestep_type is None:
|
||||
timestep_type = self.train_config.timestep_type
|
||||
|
||||
# make fake latents
|
||||
l = torch.randn(
|
||||
true_batch_size, 16, height, width
|
||||
).to(self.device_torch, dtype=dtype)
|
||||
|
||||
self.sd.noise_scheduler.set_train_timesteps(
|
||||
self.train_config.max_denoising_steps,
|
||||
device=self.device_torch,
|
||||
timestep_type=timestep_type,
|
||||
latents=l
|
||||
)
|
||||
else:
|
||||
self.sd.noise_scheduler.set_timesteps(
|
||||
self.train_config.max_denoising_steps, device=self.device_torch
|
||||
)
|
||||
|
||||
# ger a random number of steps
|
||||
timesteps_to = torch.randint(
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
torch
|
||||
torchvision
|
||||
torch==2.6.0
|
||||
torchvision==0.21.0
|
||||
torchao==0.9.0
|
||||
safetensors
|
||||
git+https://github.com/huggingface/diffusers.git
|
||||
transformers
|
||||
git+https://github.com/huggingface/diffusers@363d1ab7e24c5ed6c190abb00df66d9edb74383b
|
||||
transformers==4.49.0
|
||||
lycoris-lora==1.8.3
|
||||
flatten_json
|
||||
pyyaml
|
||||
@@ -13,7 +14,8 @@ invisible-watermark
|
||||
einops
|
||||
accelerate
|
||||
toml
|
||||
albumentations
|
||||
albumentations==1.4.15
|
||||
albucore==0.0.16
|
||||
pydantic
|
||||
omegaconf
|
||||
k-diffusion
|
||||
@@ -26,7 +28,11 @@ bitsandbytes
|
||||
hf_transfer
|
||||
lpips
|
||||
pytorch_fid
|
||||
optimum-quanto
|
||||
optimum-quanto==0.2.4
|
||||
sentencepiece
|
||||
huggingface_hub
|
||||
peft
|
||||
peft
|
||||
gradio
|
||||
python-slugify
|
||||
opencv-python
|
||||
pytorch-wavelets==1.3.0
|
||||
45
run.py
45
run.py
@@ -20,20 +20,26 @@ if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
from toolkit.accelerator import get_accelerator
|
||||
from toolkit.print import print_acc, setup_log_to_file
|
||||
|
||||
accelerator = get_accelerator()
|
||||
|
||||
|
||||
def print_end_message(jobs_completed, jobs_failed):
|
||||
if not accelerator.is_main_process:
|
||||
return
|
||||
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:")
|
||||
print_acc("")
|
||||
print_acc("========================================")
|
||||
print_acc("Result:")
|
||||
if len(completed_string) > 0:
|
||||
print(f" - {completed_string}")
|
||||
print_acc(f" - {completed_string}")
|
||||
if len(failure_string) > 0:
|
||||
print(f" - {failure_string}")
|
||||
print("========================================")
|
||||
print_acc(f" - {failure_string}")
|
||||
print_acc("========================================")
|
||||
|
||||
|
||||
def main():
|
||||
@@ -61,7 +67,17 @@ def main():
|
||||
default=None,
|
||||
help='Name to replace [name] tag in config file, useful for shared config file'
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
'-l', '--log',
|
||||
type=str,
|
||||
default=None,
|
||||
help='Log file to write output to'
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.log is not None:
|
||||
setup_log_to_file(args.log)
|
||||
|
||||
config_file_list = args.config_file_list
|
||||
if len(config_file_list) == 0:
|
||||
@@ -70,7 +86,8 @@ def main():
|
||||
jobs_completed = 0
|
||||
jobs_failed = 0
|
||||
|
||||
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
if accelerator.is_main_process:
|
||||
print_acc(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
|
||||
|
||||
for config_file in config_file_list:
|
||||
try:
|
||||
@@ -79,8 +96,20 @@ def main():
|
||||
job.cleanup()
|
||||
jobs_completed += 1
|
||||
except Exception as e:
|
||||
print(f"Error running job: {e}")
|
||||
print_acc(f"Error running job: {e}")
|
||||
jobs_failed += 1
|
||||
try:
|
||||
job.process[0].on_error(e)
|
||||
except Exception as e2:
|
||||
print_acc(f"Error running on_error: {e2}")
|
||||
if not args.recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
except KeyboardInterrupt as e:
|
||||
try:
|
||||
job.process[0].on_error(e)
|
||||
except Exception as e2:
|
||||
print_acc(f"Error running on_error: {e2}")
|
||||
if not args.recover:
|
||||
print_end_message(jobs_completed, jobs_failed)
|
||||
raise e
|
||||
|
||||
426
scripts/convert_diffusers_to_comfy.py
Normal file
426
scripts/convert_diffusers_to_comfy.py
Normal file
@@ -0,0 +1,426 @@
|
||||
#######################################################
|
||||
# Convert Diffusers Flux/Flex to all in one ComfyUI safetensors file
|
||||
# The VAE, T5 and clip will all be in the safetensors file
|
||||
# T5 will always be 8bit with the all in one file
|
||||
# You can save the transformer weights as bf16 or 8-bit with the --do_8_bit flag
|
||||
#
|
||||
# Download a reference model from Huggingface
|
||||
# https://huggingface.co/Comfy-Org/flux1-dev/blob/main/flux1-dev-fp8.safetensors
|
||||
#
|
||||
# Call like this for 8-bit transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors --do_8_bit
|
||||
#
|
||||
# Call like this for bf16 transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors
|
||||
#
|
||||
#######################################################
|
||||
|
||||
|
||||
import argparse
|
||||
from datetime import date
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import safetensors
|
||||
import safetensors.torch
|
||||
import torch
|
||||
import tqdm
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("diffusers_path", type=str,
|
||||
help="Path to the original Flux diffusers folder.")
|
||||
parser.add_argument("quantized_state_dict_path", type=str,
|
||||
help="Path to the ComfyUI all in one template file.")
|
||||
parser.add_argument("flux_path", type=str,
|
||||
help="Output path for the Flux safetensors file.")
|
||||
parser.add_argument("--do_8_bit", action="store_true",
|
||||
help="Use 8-bit weights instead of bf16.")
|
||||
args = parser.parse_args()
|
||||
|
||||
flux_path = Path(args.flux_path)
|
||||
diffusers_path = Path(args.diffusers_path, "transformer")
|
||||
quantized_state_dict_path = Path(args.quantized_state_dict_path)
|
||||
|
||||
do_8_bit = args.do_8_bit
|
||||
|
||||
if not os.path.exists(flux_path.parent):
|
||||
os.makedirs(flux_path.parent)
|
||||
|
||||
if not diffusers_path.exists():
|
||||
print(f"Error: Missing transformer folder: {diffusers_path}")
|
||||
exit()
|
||||
|
||||
original_json_path = Path.joinpath(
|
||||
diffusers_path, "diffusion_pytorch_model.safetensors.index.json")
|
||||
if not original_json_path.exists():
|
||||
print(f"Error: Missing transformer index json: {original_json_path}")
|
||||
exit()
|
||||
|
||||
if not os.path.exists(quantized_state_dict_path):
|
||||
print(
|
||||
f"Error: Missing quantized state dict file: {args.quantized_state_dict_path}")
|
||||
exit()
|
||||
|
||||
with open(original_json_path, "r", encoding="utf-8") as f:
|
||||
original_json = json.load(f)
|
||||
|
||||
diffusers_map = {
|
||||
"time_in.in_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_1.weight",
|
||||
],
|
||||
"time_in.in_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_1.bias",
|
||||
],
|
||||
"time_in.out_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_2.weight",
|
||||
],
|
||||
"time_in.out_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_2.bias",
|
||||
],
|
||||
"vector_in.in_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_1.weight",
|
||||
],
|
||||
"vector_in.in_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_1.bias",
|
||||
],
|
||||
"vector_in.out_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_2.weight",
|
||||
],
|
||||
"vector_in.out_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_2.bias",
|
||||
],
|
||||
"guidance_in.in_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_1.weight",
|
||||
],
|
||||
"guidance_in.in_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_1.bias",
|
||||
],
|
||||
"guidance_in.out_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_2.weight",
|
||||
],
|
||||
"guidance_in.out_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_2.bias",
|
||||
],
|
||||
"txt_in.weight": [
|
||||
"context_embedder.weight",
|
||||
],
|
||||
"txt_in.bias": [
|
||||
"context_embedder.bias",
|
||||
],
|
||||
"img_in.weight": [
|
||||
"x_embedder.weight",
|
||||
],
|
||||
"img_in.bias": [
|
||||
"x_embedder.bias",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.weight": [
|
||||
"norm1.linear.weight",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.bias": [
|
||||
"norm1.linear.bias",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.weight": [
|
||||
"norm1_context.linear.weight",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.bias": [
|
||||
"norm1_context.linear.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.weight": [
|
||||
"attn.add_q_proj.weight",
|
||||
"attn.add_k_proj.weight",
|
||||
"attn.add_v_proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.bias": [
|
||||
"attn.add_q_proj.bias",
|
||||
"attn.add_k_proj.bias",
|
||||
"attn.add_v_proj.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.query_norm.scale": [
|
||||
"attn.norm_added_q.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.key_norm.scale": [
|
||||
"attn.norm_added_k.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.weight": [
|
||||
"ff.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.bias": [
|
||||
"ff.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.weight": [
|
||||
"ff.net.2.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.bias": [
|
||||
"ff.net.2.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.weight": [
|
||||
"ff_context.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.bias": [
|
||||
"ff_context.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.weight": [
|
||||
"ff_context.net.2.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.bias": [
|
||||
"ff_context.net.2.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.weight": [
|
||||
"attn.to_out.0.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.bias": [
|
||||
"attn.to_out.0.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.weight": [
|
||||
"attn.to_add_out.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.bias": [
|
||||
"attn.to_add_out.bias",
|
||||
],
|
||||
"single_blocks.().modulation.lin.weight": [
|
||||
"norm.linear.weight",
|
||||
],
|
||||
"single_blocks.().modulation.lin.bias": [
|
||||
"norm.linear.bias",
|
||||
],
|
||||
"single_blocks.().linear1.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
"proj_mlp.weight",
|
||||
],
|
||||
"single_blocks.().linear1.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
"proj_mlp.bias",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"single_blocks.().norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().linear2.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.linear.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"final_layer.linear.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.weight": [
|
||||
"norm_out.linear.weight",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.bias": [
|
||||
"norm_out.linear.bias",
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def is_in_diffusers_map(k):
|
||||
for values in diffusers_map.values():
|
||||
for value in values:
|
||||
if k.endswith(value):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
diffusers = {k: Path.joinpath(diffusers_path, v)
|
||||
for k, v in original_json["weight_map"].items() if is_in_diffusers_map(k)}
|
||||
|
||||
original_safetensors = set(diffusers.values())
|
||||
|
||||
# determine the number of transformer blocks
|
||||
transformer_blocks = 0
|
||||
single_transformer_blocks = 0
|
||||
for key in diffusers.keys():
|
||||
print(key)
|
||||
if key.startswith("transformer_blocks."):
|
||||
print(key)
|
||||
block = int(key.split(".")[1])
|
||||
if block >= transformer_blocks:
|
||||
transformer_blocks = block + 1
|
||||
elif key.startswith("single_transformer_blocks."):
|
||||
block = int(key.split(".")[1])
|
||||
if block >= single_transformer_blocks:
|
||||
single_transformer_blocks = block + 1
|
||||
|
||||
print(f"Transformer blocks: {transformer_blocks}")
|
||||
print(f"Single transformer blocks: {single_transformer_blocks}")
|
||||
|
||||
for file in original_safetensors:
|
||||
if not file.exists():
|
||||
print(f"Error: Missing transformer safetensors file: {file}")
|
||||
exit()
|
||||
|
||||
original_safetensors = {f: safetensors.safe_open(
|
||||
f, framework="pt", device="cpu") for f in original_safetensors}
|
||||
|
||||
|
||||
def swap_scale_shift(weight):
|
||||
shift, scale = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([scale, shift], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
flux_values = {}
|
||||
|
||||
for b in range(transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("double_blocks."):
|
||||
block_prefix = f"transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
for b in range(single_transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("single_blocks."):
|
||||
block_prefix = f"single_transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
|
||||
for key, weights in diffusers_map.items():
|
||||
if not (key.startswith("double_blocks.") or key.startswith("single_blocks.")):
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key] = [f"{weight}" for weight in weights]
|
||||
|
||||
flux = {}
|
||||
|
||||
for key, values in tqdm.tqdm(flux_values.items()):
|
||||
if len(values) == 1:
|
||||
flux[key] = original_safetensors[diffusers[values[0]]
|
||||
].get_tensor(values[0]).to("cpu")
|
||||
else:
|
||||
flux[key] = torch.cat(
|
||||
[
|
||||
original_safetensors[diffusers[value]
|
||||
].get_tensor(value).to("cpu")
|
||||
for value in values
|
||||
]
|
||||
)
|
||||
|
||||
if "norm_out.linear.weight" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.weight"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.weight"]].get_tensor(
|
||||
"norm_out.linear.weight").to("cpu")
|
||||
)
|
||||
if "norm_out.linear.bias" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.bias"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.bias"]].get_tensor(
|
||||
"norm_out.linear.bias").to("cpu")
|
||||
)
|
||||
|
||||
|
||||
def stochastic_round_to(tensor, dtype=torch.float8_e4m3fn):
|
||||
# Define the float8 range
|
||||
min_val = torch.finfo(dtype).min
|
||||
max_val = torch.finfo(dtype).max
|
||||
|
||||
# Clip values to float8 range
|
||||
tensor = torch.clamp(tensor, min_val, max_val)
|
||||
|
||||
# Convert to float32 for calculations
|
||||
tensor = tensor.float()
|
||||
|
||||
# Get the nearest representable float8 values
|
||||
lower = torch.floor(tensor * 256) / 256
|
||||
upper = torch.ceil(tensor * 256) / 256
|
||||
|
||||
# Calculate the probability of rounding up
|
||||
prob = (tensor - lower) / (upper - lower)
|
||||
|
||||
# Generate random values for stochastic rounding
|
||||
rand = torch.rand_like(tensor)
|
||||
|
||||
# Perform stochastic rounding
|
||||
rounded = torch.where(rand < prob, upper, lower)
|
||||
|
||||
# Convert back to float8
|
||||
return rounded.to(dtype)
|
||||
|
||||
|
||||
# set all the keys to bf16
|
||||
for key in flux.keys():
|
||||
if do_8_bit:
|
||||
flux[key] = stochastic_round_to(
|
||||
flux[key], torch.float8_e4m3fn).to('cpu')
|
||||
else:
|
||||
flux[key] = flux[key].clone().to('cpu', torch.bfloat16)
|
||||
|
||||
# load the quantized state dict
|
||||
quantized_state_dict = safetensors.torch.load_file(quantized_state_dict_path)
|
||||
|
||||
transformer_pre = "model.diffusion_model."
|
||||
did_print = False
|
||||
# remove old parts
|
||||
for key in list(quantized_state_dict.keys()):
|
||||
if key.startswith(transformer_pre):
|
||||
if not did_print:
|
||||
# print("dtype: ", quantized_state_dict[key].dtype)
|
||||
did_print = True
|
||||
del quantized_state_dict[key]
|
||||
|
||||
# add the new parts
|
||||
for key, value in flux.items():
|
||||
quantized_state_dict[transformer_pre + key] = value
|
||||
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
# date format like 2024-08-01 YYYY-MM-DD
|
||||
meta['modelspec.date'] = date.today().strftime("%Y-%m-%d")
|
||||
meta['modelspec.title'] = "Flex.1-alpha"
|
||||
meta['modelspec.author'] = "Ostris, LLC"
|
||||
meta['modelspec.license'] = "Apache-2.0"
|
||||
meta['modelspec.implementation'] = "https://github.com/black-forest-labs/flux"
|
||||
meta['modelspec.architecture'] = "Flex.1-alpha"
|
||||
meta['modelspec.description'] = "Flex.1-alpha"
|
||||
|
||||
|
||||
os.makedirs(os.path.dirname(flux_path), exist_ok=True)
|
||||
|
||||
print(f"Saving to {flux_path}")
|
||||
|
||||
safetensors.torch.save_file(quantized_state_dict, flux_path, metadata=meta)
|
||||
|
||||
print("Done.")
|
||||
245
scripts/extract_lora_from_flex.py
Normal file
245
scripts/extract_lora_from_flex.py
Normal file
@@ -0,0 +1,245 @@
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
|
||||
parser = argparse.ArgumentParser(description="Extract LoRA from Flex")
|
||||
parser.add_argument("--base", type=str, default="ostris/Flex.1-alpha", help="Base model path")
|
||||
parser.add_argument("--tuned", type=str, required=True, help="Tuned model path")
|
||||
parser.add_argument("--output", type=str, required=True, help="Output path for lora")
|
||||
parser.add_argument("--rank", type=int, default=32, help="LoRA rank for extraction")
|
||||
parser.add_argument("--gpu", type=int, default=0, help="GPU to process extraction")
|
||||
parser.add_argument("--full", action="store_true", help="Do a full transformer extraction, not just transformer blocks")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if True:
|
||||
# set cuda environment variable
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from lycoris.utils import extract_linear, extract_conv, make_sparse
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
base = args.base
|
||||
tuned = args.tuned
|
||||
output_path = args.output
|
||||
dim = args.rank
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
|
||||
state_dict_base = {}
|
||||
state_dict_tuned = {}
|
||||
|
||||
output_dict = {}
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_diff(
|
||||
base_unet,
|
||||
db_unet,
|
||||
mode="fixed",
|
||||
linear_mode_param=0,
|
||||
conv_mode_param=0,
|
||||
extract_device="cpu",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
# small_conv=True,
|
||||
small_conv=False,
|
||||
):
|
||||
UNET_TARGET_REPLACE_MODULE = [
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
]
|
||||
LORA_PREFIX_UNET = "transformer"
|
||||
|
||||
def make_state_dict(
|
||||
prefix,
|
||||
root_module: torch.nn.Module,
|
||||
target_module: torch.nn.Module,
|
||||
target_replace_modules,
|
||||
):
|
||||
loras = {}
|
||||
temp = {}
|
||||
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
temp[name] = module
|
||||
|
||||
for name, module in tqdm(
|
||||
list((n, m) for n, m in target_module.named_modules() if n in temp)
|
||||
):
|
||||
weights = temp[name]
|
||||
lora_name = prefix + "." + name
|
||||
# lora_name = lora_name.replace(".", "_")
|
||||
layer = module.__class__.__name__
|
||||
if 'transformer_blocks' not in lora_name and not args.full:
|
||||
continue
|
||||
|
||||
if layer in {
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"Embedding",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
}:
|
||||
root_weight = module.weight
|
||||
try:
|
||||
if torch.allclose(root_weight, weights.weight):
|
||||
continue
|
||||
except:
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
module = module.to(extract_device, torch.float32)
|
||||
weights = weights.to(extract_device, torch.float32)
|
||||
|
||||
if mode == "full":
|
||||
decompose_mode = "full"
|
||||
elif layer == "Linear":
|
||||
weight, decompose_mode = extract_linear(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
elif layer == "Conv2d":
|
||||
is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1
|
||||
weight, decompose_mode = extract_conv(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param if is_linear else conv_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
if small_conv and not is_linear and decompose_mode == "low rank":
|
||||
dim = extract_a.size(0)
|
||||
(extract_c, extract_a, _), _ = extract_conv(
|
||||
extract_a.transpose(0, 1),
|
||||
"fixed",
|
||||
dim,
|
||||
extract_device,
|
||||
True,
|
||||
)
|
||||
extract_a = extract_a.transpose(0, 1)
|
||||
extract_c = extract_c.transpose(0, 1)
|
||||
loras[f"{lora_name}.lora_mid.weight"] = (
|
||||
extract_c.detach().cpu().contiguous().half()
|
||||
)
|
||||
diff = (
|
||||
(
|
||||
root_weight
|
||||
- torch.einsum(
|
||||
"i j k l, j r, p i -> p r k l",
|
||||
extract_c,
|
||||
extract_a.flatten(1, -1),
|
||||
extract_b.flatten(1, -1),
|
||||
)
|
||||
)
|
||||
.detach()
|
||||
.cpu()
|
||||
.contiguous()
|
||||
)
|
||||
del extract_c
|
||||
else:
|
||||
module = module.to("cpu")
|
||||
weights = weights.to("cpu")
|
||||
continue
|
||||
|
||||
if decompose_mode == "low rank":
|
||||
loras[f"{lora_name}.lora_A.weight"] = (
|
||||
extract_a.detach().cpu().contiguous().half()
|
||||
)
|
||||
loras[f"{lora_name}.lora_B.weight"] = (
|
||||
extract_b.detach().cpu().contiguous().half()
|
||||
)
|
||||
# loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half()
|
||||
if use_bias:
|
||||
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
|
||||
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
|
||||
|
||||
indices = sparse_diff.indices().to(torch.int16)
|
||||
values = sparse_diff.values().half()
|
||||
loras[f"{lora_name}.bias_indices"] = indices
|
||||
loras[f"{lora_name}.bias_values"] = values
|
||||
loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to(
|
||||
torch.int16
|
||||
)
|
||||
del extract_a, extract_b, diff
|
||||
elif decompose_mode == "full":
|
||||
if "Norm" in layer:
|
||||
w_key = "w_norm"
|
||||
b_key = "b_norm"
|
||||
else:
|
||||
w_key = "diff"
|
||||
b_key = "diff_b"
|
||||
weight_diff = module.weight - weights.weight
|
||||
loras[f"{lora_name}.{w_key}"] = (
|
||||
weight_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
if getattr(weights, "bias", None) is not None:
|
||||
bias_diff = module.bias - weights.bias
|
||||
loras[f"{lora_name}.{b_key}"] = (
|
||||
bias_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
module = module.to("cpu", torch.bfloat16)
|
||||
weights = weights.to("cpu", torch.bfloat16)
|
||||
return loras
|
||||
|
||||
all_loras = {}
|
||||
|
||||
all_loras |= make_state_dict(
|
||||
LORA_PREFIX_UNET,
|
||||
base_unet,
|
||||
db_unet,
|
||||
UNET_TARGET_REPLACE_MODULE,
|
||||
)
|
||||
del base_unet, db_unet
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_lora_name = set()
|
||||
for k in all_loras:
|
||||
lora_name, weight = k.rsplit(".", 1)
|
||||
all_lora_name.add(lora_name)
|
||||
print(len(all_lora_name))
|
||||
return all_loras
|
||||
|
||||
|
||||
# find all the .safetensors files and load them
|
||||
print("Loading Base")
|
||||
base_model = FluxTransformer2DModel.from_pretrained(base, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
print("Loading Tuned")
|
||||
tuned_model = FluxTransformer2DModel.from_pretrained(tuned, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
output_dict = extract_diff(
|
||||
base_model,
|
||||
tuned_model,
|
||||
mode="fixed",
|
||||
linear_mode_param=dim,
|
||||
conv_mode_param=dim,
|
||||
extract_device="cuda",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
small_conv=False,
|
||||
)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
|
||||
save_file(output_dict, output_path, metadata=meta)
|
||||
|
||||
print("Done")
|
||||
309
scripts/update_sponsors.py
Normal file
309
scripts/update_sponsors.py
Normal file
@@ -0,0 +1,309 @@
|
||||
import os
|
||||
import requests
|
||||
import json
|
||||
from datetime import datetime
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Load environment variables from .env file
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(__file__)), ".env")
|
||||
load_dotenv(dotenv_path=env_path)
|
||||
|
||||
# API credentials
|
||||
PATREON_TOKEN = os.getenv("PATREON_ACCESS_TOKEN")
|
||||
GITHUB_TOKEN = os.getenv("GITHUB_TOKEN")
|
||||
GITHUB_USERNAME = os.getenv("GITHUB_USERNAME")
|
||||
GITHUB_ORG = os.getenv("GITHUB_ORG") # Organization name (optional)
|
||||
|
||||
# Output file
|
||||
README_PATH = "SUPPORTERS.md"
|
||||
|
||||
def fetch_patreon_supporters():
|
||||
"""Fetch current Patreon supporters"""
|
||||
print("Fetching Patreon supporters...")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {PATREON_TOKEN}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
|
||||
url = "https://www.patreon.com/api/oauth2/v2/campaigns"
|
||||
|
||||
try:
|
||||
# First get the campaign ID
|
||||
campaign_response = requests.get(url, headers=headers)
|
||||
campaign_response.raise_for_status()
|
||||
campaign_data = campaign_response.json()
|
||||
|
||||
if not campaign_data.get('data'):
|
||||
print("No campaigns found for this Patreon account")
|
||||
return []
|
||||
|
||||
campaign_id = campaign_data['data'][0]['id']
|
||||
|
||||
# Now get the supporters for this campaign
|
||||
members_url = f"https://www.patreon.com/api/oauth2/v2/campaigns/{campaign_id}/members"
|
||||
params = {
|
||||
"include": "user",
|
||||
"fields[member]": "full_name,is_follower,patron_status", # Removed profile_url
|
||||
"fields[user]": "image_url"
|
||||
}
|
||||
|
||||
supporters = []
|
||||
while members_url:
|
||||
members_response = requests.get(members_url, headers=headers, params=params)
|
||||
members_response.raise_for_status()
|
||||
members_data = members_response.json()
|
||||
|
||||
# Process the response to extract active patrons
|
||||
for member in members_data.get('data', []):
|
||||
attributes = member.get('attributes', {})
|
||||
|
||||
# Only include active patrons
|
||||
if attributes.get('patron_status') == 'active_patron':
|
||||
name = attributes.get('full_name', 'Anonymous Supporter')
|
||||
|
||||
# Get user data which contains the profile image
|
||||
user_id = member.get('relationships', {}).get('user', {}).get('data', {}).get('id')
|
||||
profile_image = None
|
||||
profile_url = None # Removed profile_url since it's not supported
|
||||
|
||||
if user_id:
|
||||
for included in members_data.get('included', []):
|
||||
if included.get('id') == user_id and included.get('type') == 'user':
|
||||
profile_image = included.get('attributes', {}).get('image_url')
|
||||
break
|
||||
|
||||
supporters.append({
|
||||
'name': name,
|
||||
'profile_image': profile_image,
|
||||
'profile_url': profile_url, # This will be None
|
||||
'platform': 'Patreon',
|
||||
'amount': 0 # Placeholder, as Patreon API doesn't provide this in the current response
|
||||
})
|
||||
|
||||
# Handle pagination
|
||||
members_url = members_data.get('links', {}).get('next')
|
||||
|
||||
print(f"Found {len(supporters)} active Patreon supporters")
|
||||
return supporters
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error fetching Patreon data: {e}")
|
||||
print(f"Response content: {e.response.content if hasattr(e, 'response') else 'No response content'}")
|
||||
return []
|
||||
|
||||
def fetch_github_sponsors():
|
||||
"""Fetch current GitHub sponsors for a user or organization"""
|
||||
print("Fetching GitHub sponsors...")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {GITHUB_TOKEN}",
|
||||
"Accept": "application/vnd.github.v3+json"
|
||||
}
|
||||
|
||||
# Determine if we're fetching for a user or an organization
|
||||
entity_type = "organization" if GITHUB_ORG else "user"
|
||||
entity_name = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
|
||||
|
||||
if not entity_name:
|
||||
print("Error: Neither GITHUB_USERNAME nor GITHUB_ORG is set")
|
||||
return []
|
||||
|
||||
# Different GraphQL query structure based on entity type
|
||||
if entity_type == "user":
|
||||
query = """
|
||||
query {
|
||||
user(login: "%s") {
|
||||
sponsorshipsAsMaintainer(first: 100) {
|
||||
nodes {
|
||||
sponsorEntity {
|
||||
... on User {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
... on Organization {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
}
|
||||
tier {
|
||||
monthlyPriceInDollars
|
||||
}
|
||||
isOneTimePayment
|
||||
isActive
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
""" % entity_name
|
||||
else: # organization
|
||||
query = """
|
||||
query {
|
||||
organization(login: "%s") {
|
||||
sponsorshipsAsMaintainer(first: 100) {
|
||||
nodes {
|
||||
sponsorEntity {
|
||||
... on User {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
... on Organization {
|
||||
login
|
||||
name
|
||||
avatarUrl
|
||||
url
|
||||
}
|
||||
}
|
||||
tier {
|
||||
monthlyPriceInDollars
|
||||
}
|
||||
isOneTimePayment
|
||||
isActive
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
""" % entity_name
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
"https://api.github.com/graphql",
|
||||
headers=headers,
|
||||
json={"query": query}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
# Process the response - the path to the data differs based on entity type
|
||||
if entity_type == "user":
|
||||
sponsors_data = data.get('data', {}).get('user', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
|
||||
else:
|
||||
sponsors_data = data.get('data', {}).get('organization', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
|
||||
|
||||
sponsors = []
|
||||
for sponsor in sponsors_data:
|
||||
# Only include active sponsors
|
||||
if sponsor.get('isActive'):
|
||||
entity = sponsor.get('sponsorEntity', {})
|
||||
name = entity.get('name') or entity.get('login', 'Anonymous Sponsor')
|
||||
profile_image = entity.get('avatarUrl')
|
||||
profile_url = entity.get('url')
|
||||
amount = sponsor.get('tier', {}).get('monthlyPriceInDollars', 0)
|
||||
|
||||
sponsors.append({
|
||||
'name': name,
|
||||
'profile_image': profile_image,
|
||||
'profile_url': profile_url,
|
||||
'platform': 'GitHub Sponsors',
|
||||
'amount': amount
|
||||
})
|
||||
|
||||
print(f"Found {len(sponsors)} active GitHub sponsors for {entity_type} '{entity_name}'")
|
||||
return sponsors
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error fetching GitHub sponsors data: {e}")
|
||||
return []
|
||||
|
||||
def generate_readme(supporters):
|
||||
"""Generate a README.md file with supporter information"""
|
||||
print(f"Generating {README_PATH}...")
|
||||
|
||||
# Sort supporters by amount (descending) and then by name
|
||||
supporters.sort(key=lambda x: (-x['amount'], x['name'].lower()))
|
||||
|
||||
# Determine the proper footer links based on what's configured
|
||||
github_entity = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
|
||||
github_entity_type = "orgs" if GITHUB_ORG else "sponsors"
|
||||
github_sponsor_url = f"https://github.com/{github_entity_type}/{github_entity}"
|
||||
|
||||
with open(README_PATH, "w", encoding="utf-8") as f:
|
||||
f.write("## Support My Work\n\n")
|
||||
f.write("If you enjoy my work, or use it for commercial purposes, please consider sponsoring me so I can continue to maintain it. Every bit helps! \n\n")
|
||||
# Create appropriate call-to-action based on what's configured
|
||||
cta_parts = []
|
||||
if github_entity:
|
||||
cta_parts.append(f"[Become a sponsor on GitHub]({github_sponsor_url})")
|
||||
if PATREON_TOKEN:
|
||||
cta_parts.append("[support me on Patreon](https://www.patreon.com/ostris)")
|
||||
|
||||
if cta_parts:
|
||||
if GITHUB_ORG:
|
||||
f.write(f"{' or '.join(cta_parts)}.\n\n")
|
||||
f.write("Thank you to all my current supporters!\n\n")
|
||||
|
||||
f.write(f"_Last updated: {datetime.now().strftime('%Y-%m-%d')}_\n\n")
|
||||
|
||||
# Write GitHub Sponsors section
|
||||
github_sponsors = [s for s in supporters if s['platform'] == 'GitHub Sponsors']
|
||||
if github_sponsors:
|
||||
f.write("### GitHub Sponsors\n\n")
|
||||
for sponsor in github_sponsors:
|
||||
if sponsor['profile_image']:
|
||||
f.write(f"<a href=\"{sponsor['profile_url']}\" title=\"{sponsor['name']}\"><img src=\"{sponsor['profile_image']}\" width=\"50\" height=\"50\" alt=\"{sponsor['name']}\" style=\"border-radius:50%\"></a> ")
|
||||
else:
|
||||
f.write(f"[{sponsor['name']}]({sponsor['profile_url']}) ")
|
||||
f.write("\n\n")
|
||||
|
||||
# Write Patreon section
|
||||
patreon_supporters = [s for s in supporters if s['platform'] == 'Patreon']
|
||||
if patreon_supporters:
|
||||
f.write("### Patreon Supporters\n\n")
|
||||
for supporter in patreon_supporters:
|
||||
if supporter['profile_image']:
|
||||
f.write(f"<a href=\"{supporter['profile_url']}\" title=\"{supporter['name']}\"><img src=\"{supporter['profile_image']}\" width=\"50\" height=\"50\" alt=\"{supporter['name']}\" style=\"border-radius:50%\"></a> ")
|
||||
else:
|
||||
f.write(f"[{supporter['name']}]({supporter['profile_url']}) ")
|
||||
f.write("\n\n")
|
||||
|
||||
f.write("\n---\n\n")
|
||||
|
||||
|
||||
print(f"Successfully generated {README_PATH} with {len(supporters)} supporters!")
|
||||
|
||||
def main():
|
||||
"""Main function"""
|
||||
print("Starting supporter data collection...")
|
||||
|
||||
# Check if required environment variables are set
|
||||
missing_vars = []
|
||||
if not GITHUB_TOKEN:
|
||||
missing_vars.append("GITHUB_TOKEN")
|
||||
|
||||
# Either username or org is required for GitHub
|
||||
if not GITHUB_USERNAME and not GITHUB_ORG:
|
||||
missing_vars.append("GITHUB_USERNAME or GITHUB_ORG")
|
||||
|
||||
# Patreon token is optional but warn if missing
|
||||
patreon_enabled = bool(PATREON_TOKEN)
|
||||
|
||||
if missing_vars:
|
||||
print(f"Error: Missing required environment variables: {', '.join(missing_vars)}")
|
||||
print("Please add them to your .env file")
|
||||
return
|
||||
|
||||
if not patreon_enabled:
|
||||
print("Warning: PATREON_ACCESS_TOKEN not set. Will only fetch GitHub sponsors.")
|
||||
|
||||
# Fetch data from both platforms
|
||||
patreon_supporters = fetch_patreon_supporters() if PATREON_TOKEN else []
|
||||
github_sponsors = fetch_github_sponsors()
|
||||
|
||||
# Combine supporters from both platforms
|
||||
all_supporters = patreon_supporters + github_sponsors
|
||||
|
||||
if not all_supporters:
|
||||
print("No supporters found on either platform")
|
||||
return
|
||||
|
||||
# Generate README
|
||||
generate_readme(all_supporters)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -13,7 +13,7 @@ from transformers import CLIPImageProcessor
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from toolkit.paths import SD_SCRIPTS_ROOT
|
||||
import torchvision.transforms.functional
|
||||
from toolkit.image_utils import show_img, show_tensors
|
||||
from toolkit.image_utils import save_tensors, show_img, show_tensors
|
||||
|
||||
sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
@@ -28,13 +28,18 @@ from tqdm import tqdm
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('dataset_folder', type=str, default='input')
|
||||
parser.add_argument('--epochs', type=int, default=1)
|
||||
|
||||
parser.add_argument('--num_frames', type=int, default=1)
|
||||
parser.add_argument('--output_path', type=str, default=None)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.output_path is not None:
|
||||
args.output_path = os.path.abspath(args.output_path)
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
dataset_folder = args.dataset_folder
|
||||
resolution = 1024
|
||||
resolution = 512
|
||||
bucket_tolerance = 64
|
||||
batch_size = 1
|
||||
|
||||
@@ -63,6 +68,8 @@ dataset_config = DatasetConfig(
|
||||
# clip_image_path='/mnt/Datasets2/regs/yetibear_xl_v14/random_aspect/',
|
||||
buckets=True,
|
||||
bucket_tolerance=bucket_tolerance,
|
||||
shrink_video_to_frames=True,
|
||||
num_frames=args.num_frames,
|
||||
# poi='person',
|
||||
# shuffle_augmentations=True,
|
||||
# augmentations=[
|
||||
@@ -80,11 +87,17 @@ dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_si
|
||||
|
||||
# run through an epoch ang check sizes
|
||||
dataloader_iterator = iter(dataloader)
|
||||
idx = 0
|
||||
for epoch in range(args.epochs):
|
||||
for batch in tqdm(dataloader):
|
||||
batch: 'DataLoaderBatchDTO'
|
||||
img_batch = batch.tensor
|
||||
batch_size, channels, height, width = img_batch.shape
|
||||
frames = 1
|
||||
if len(img_batch.shape) == 5:
|
||||
frames = img_batch.shape[1]
|
||||
batch_size, frames, channels, height, width = img_batch.shape
|
||||
else:
|
||||
batch_size, channels, height, width = img_batch.shape
|
||||
|
||||
# img_batch = color_block_imgs(img_batch, neg1_1=True)
|
||||
|
||||
@@ -110,15 +123,18 @@ for epoch in range(args.epochs):
|
||||
|
||||
big_img = img_batch
|
||||
# big_img = big_img.clamp(-1, 1)
|
||||
if args.output_path is not None:
|
||||
save_tensors(big_img, os.path.join(args.output_path, f'{idx}.png'))
|
||||
else:
|
||||
show_tensors(big_img)
|
||||
|
||||
show_tensors(big_img)
|
||||
# convert to image
|
||||
# img = transforms.ToPILImage()(big_img)
|
||||
#
|
||||
# show_img(img)
|
||||
|
||||
# convert to image
|
||||
# img = transforms.ToPILImage()(big_img)
|
||||
#
|
||||
# show_img(img)
|
||||
|
||||
time.sleep(0.2)
|
||||
time.sleep(0.2)
|
||||
idx += 1
|
||||
# if not last epoch
|
||||
if epoch < args.epochs - 1:
|
||||
trigger_dataloader_setup_epoch(dataloader)
|
||||
|
||||
@@ -29,7 +29,7 @@ def paramiter_count(model):
|
||||
return int(paramiter_count)
|
||||
|
||||
|
||||
def calculate_metrics(vae, images, max_imgs=-1):
|
||||
def calculate_metrics(vae, images, max_imgs=-1, save_output=False):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
vae = vae.to(device)
|
||||
lpips_model = lpips.LPIPS(net='alex').to(device)
|
||||
@@ -44,6 +44,9 @@ def calculate_metrics(vae, images, max_imgs=-1):
|
||||
# ])
|
||||
# needs values between -1 and 1
|
||||
to_tensor = ToTensor()
|
||||
|
||||
# remove _reconstructed.png files
|
||||
images = [img for img in images if not img.endswith("_reconstructed.png")]
|
||||
|
||||
if max_imgs > 0 and len(images) > max_imgs:
|
||||
images = images[:max_imgs]
|
||||
@@ -82,6 +85,15 @@ def calculate_metrics(vae, images, max_imgs=-1):
|
||||
avg_rfid = 0
|
||||
avg_psnr = sum(psnr_scores) / len(psnr_scores)
|
||||
avg_lpips = sum(lpips_scores) / len(lpips_scores)
|
||||
|
||||
if save_output:
|
||||
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
folder = os.path.dirname(img_path)
|
||||
save_path = os.path.join(folder, filename_no_ext + "_reconstructed.png")
|
||||
reconstructed = (reconstructed + 1) / 2
|
||||
reconstructed = reconstructed.clamp(0, 1)
|
||||
reconstructed = transforms.ToPILImage()(reconstructed[0].cpu())
|
||||
reconstructed.save(save_path)
|
||||
|
||||
return avg_rfid, avg_psnr, avg_lpips
|
||||
|
||||
@@ -91,18 +103,23 @@ def main():
|
||||
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.")
|
||||
# boolean store true
|
||||
parser.add_argument("--save_output", action="store_true", help="Save the output 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)
|
||||
try:
|
||||
vae = AutoencoderKL.from_pretrained(args.vae_path)
|
||||
except:
|
||||
vae = AutoencoderKL.from_pretrained(args.vae_path, subfolder="vae")
|
||||
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)
|
||||
avg_rfid, avg_psnr, avg_lpips = calculate_metrics(vae, images, args.max_imgs, args.save_output)
|
||||
|
||||
# print(f"Average rFID: {avg_rfid}")
|
||||
print(f"Average PSNR: {avg_psnr}")
|
||||
|
||||
17
toolkit/accelerator.py
Normal file
17
toolkit/accelerator.py
Normal file
@@ -0,0 +1,17 @@
|
||||
from accelerate import Accelerator
|
||||
from diffusers.utils.torch_utils import is_compiled_module
|
||||
|
||||
global_accelerator = None
|
||||
|
||||
|
||||
def get_accelerator() -> Accelerator:
|
||||
global global_accelerator
|
||||
if global_accelerator is None:
|
||||
global_accelerator = Accelerator()
|
||||
return global_accelerator
|
||||
|
||||
def unwrap_model(model):
|
||||
accelerator = get_accelerator()
|
||||
model = accelerator.unwrap_model(model)
|
||||
model = model._orig_mod if is_compiled_module(model) else model
|
||||
return model
|
||||
@@ -56,51 +56,6 @@ resolutions_1024: List[BucketResolution] = [
|
||||
{"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},
|
||||
]
|
||||
|
||||
|
||||
def get_bucket_sizes(resolution: int = 512, divisibility: int = 8) -> List[BucketResolution]:
|
||||
# determine scaler form 1024 to resolution
|
||||
scaler = resolution / 1024
|
||||
@@ -171,4 +126,4 @@ def get_bucket_for_image_size(
|
||||
if closest_bucket is None:
|
||||
raise ValueError("No suitable bucket found")
|
||||
|
||||
return closest_bucket
|
||||
return closest_bucket
|
||||
@@ -13,7 +13,9 @@ SaveFormat = Literal['safetensors', 'diffusers']
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.guidance import GuidanceType
|
||||
|
||||
from toolkit.logging import EmptyLogger
|
||||
else:
|
||||
EmptyLogger = None
|
||||
|
||||
class SaveConfig:
|
||||
def __init__(self, **kwargs):
|
||||
@@ -27,11 +29,13 @@ class SaveConfig:
|
||||
self.hf_repo_id: Optional[str] = kwargs.get("hf_repo_id", None)
|
||||
self.hf_private: Optional[str] = kwargs.get("hf_private", False)
|
||||
|
||||
class LogingConfig:
|
||||
class LoggingConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.log_every: int = kwargs.get('log_every', 100)
|
||||
self.verbose: bool = kwargs.get('verbose', False)
|
||||
self.use_wandb: bool = kwargs.get('use_wandb', False)
|
||||
self.project_name: str = kwargs.get('project_name', 'ai-toolkit')
|
||||
self.run_name: str = kwargs.get('run_name', None)
|
||||
|
||||
|
||||
class SampleConfig:
|
||||
@@ -53,6 +57,11 @@ class SampleConfig:
|
||||
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', [])
|
||||
self.num_frames = kwargs.get('num_frames', 1)
|
||||
self.fps: int = kwargs.get('fps', 16)
|
||||
if self.num_frames > 1 and self.ext not in ['webp']:
|
||||
print("Changing sample extention to animated webp")
|
||||
self.ext = 'webp'
|
||||
|
||||
|
||||
class LormModuleSettingsConfig:
|
||||
@@ -97,7 +106,7 @@ class LoRMConfig:
|
||||
})
|
||||
|
||||
|
||||
NetworkType = Literal['lora', 'locon', 'lorm']
|
||||
NetworkType = Literal['lora', 'locon', 'lorm', 'lokr']
|
||||
|
||||
|
||||
class NetworkConfig:
|
||||
@@ -131,9 +140,18 @@ class NetworkConfig:
|
||||
self.conv = 4
|
||||
|
||||
self.transformer_only = kwargs.get('transformer_only', True)
|
||||
|
||||
self.lokr_full_rank = kwargs.get('lokr_full_rank', False)
|
||||
if self.lokr_full_rank and self.type.lower() == 'lokr':
|
||||
self.linear = 9999999999
|
||||
self.linear_alpha = 9999999999
|
||||
self.conv = 9999999999
|
||||
self.conv_alpha = 9999999999
|
||||
# -1 automatically finds the largest factor
|
||||
self.lokr_factor = kwargs.get('lokr_factor', -1)
|
||||
|
||||
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net']
|
||||
AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net', 'control_lora']
|
||||
|
||||
CLIPLayer = Literal['penultimate_hidden_states', 'image_embeds', 'last_hidden_state']
|
||||
|
||||
@@ -147,7 +165,13 @@ class AdapterConfig:
|
||||
self.downscale_factor: int = kwargs.get('downscale_factor', 8)
|
||||
self.adapter_type: str = kwargs.get('adapter_type', 'full_adapter')
|
||||
self.image_dir: str = kwargs.get('image_dir', None)
|
||||
self.test_img_path: str = kwargs.get('test_img_path', None)
|
||||
self.test_img_path: List[str] = kwargs.get('test_img_path', None)
|
||||
if self.test_img_path is not None:
|
||||
if isinstance(self.test_img_path, str):
|
||||
self.test_img_path = self.test_img_path.split(',')
|
||||
self.test_img_path = [p.strip() for p in self.test_img_path]
|
||||
self.test_img_path = [p for p in self.test_img_path if p != '']
|
||||
|
||||
self.train: str = kwargs.get('train', False)
|
||||
self.image_encoder_path: str = kwargs.get('image_encoder_path', None)
|
||||
self.name_or_path = kwargs.get('name_or_path', None)
|
||||
@@ -202,6 +226,32 @@ class AdapterConfig:
|
||||
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)
|
||||
|
||||
self.pixtral_max_image_size: int = kwargs.get('pixtral_max_image_size', 512)
|
||||
self.pixtral_random_image_size: int = kwargs.get('pixtral_random_image_size', False)
|
||||
|
||||
self.flux_only_double: bool = kwargs.get('flux_only_double', False)
|
||||
|
||||
# train and use a conv layer to pool the embedding
|
||||
self.conv_pooling: bool = kwargs.get('conv_pooling', False)
|
||||
self.conv_pooling_stacks: int = kwargs.get('conv_pooling_stacks', 1)
|
||||
self.sparse_autoencoder_dim: Optional[int] = kwargs.get('sparse_autoencoder_dim', None)
|
||||
|
||||
# for llm adapter
|
||||
self.num_cloned_blocks: int = kwargs.get('num_cloned_blocks', 0)
|
||||
self.quantize_llm: bool = kwargs.get('quantize_llm', False)
|
||||
|
||||
# for control lora only
|
||||
lora_config: dict = kwargs.get('lora_config', None)
|
||||
if lora_config is not None:
|
||||
self.lora_config: NetworkConfig = NetworkConfig(**lora_config)
|
||||
else:
|
||||
self.lora_config = None
|
||||
self.num_control_images: int = kwargs.get('num_control_images', 1)
|
||||
# decimal for how often the control is dropped out and replaced with noise 1.0 is 100%
|
||||
self.control_image_dropout: float = kwargs.get('control_image_dropout', 0.0)
|
||||
self.has_inpainting_input: bool = kwargs.get('has_inpainting_input', False)
|
||||
self.invert_inpaint_mask_chance: float = kwargs.get('invert_inpaint_mask_chance', 0.0)
|
||||
|
||||
|
||||
class EmbeddingConfig:
|
||||
@@ -213,6 +263,11 @@ class EmbeddingConfig:
|
||||
self.trigger_class_name = kwargs.get('trigger_class_name', None) # used for inverted masked prior
|
||||
|
||||
|
||||
class DecoratorConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.num_tokens: str = kwargs.get('num_tokens', 4)
|
||||
|
||||
|
||||
ContentOrStyleType = Literal['balanced', 'style', 'content']
|
||||
LossTarget = Literal['noise', 'source', 'unaugmented', 'differential_noise']
|
||||
|
||||
@@ -236,6 +291,7 @@ class TrainConfig:
|
||||
self.min_denoising_steps: int = kwargs.get('min_denoising_steps', 0)
|
||||
self.max_denoising_steps: int = kwargs.get('max_denoising_steps', 1000)
|
||||
self.batch_size: int = kwargs.get('batch_size', 1)
|
||||
self.orig_batch_size: int = self.batch_size
|
||||
self.dtype: str = kwargs.get('dtype', 'fp32')
|
||||
self.xformers = kwargs.get('xformers', False)
|
||||
self.sdp = kwargs.get('sdp', False)
|
||||
@@ -284,8 +340,16 @@ class TrainConfig:
|
||||
|
||||
# set to -1 to accumulate gradients for entire epoch
|
||||
# warning, only do this with a small dataset or you will run out of memory
|
||||
# This is legacy but left in for backwards compatibility
|
||||
self.gradient_accumulation_steps = kwargs.get('gradient_accumulation_steps', 1)
|
||||
|
||||
# this will do proper gradient accumulation where you will not see a step until the end of the accumulation
|
||||
# the method above will show a step every accumulation
|
||||
self.gradient_accumulation = kwargs.get('gradient_accumulation', 1)
|
||||
if self.gradient_accumulation > 1:
|
||||
if self.gradient_accumulation_steps != 1:
|
||||
raise ValueError("gradient_accumulation and gradient_accumulation_steps are mutually exclusive")
|
||||
|
||||
# short long captions will double your batch size. This only works when a dataset is
|
||||
# prepared with a json caption file that has both short and long captions in it. It will
|
||||
# Double up every image and run it through with both short and long captions. The idea
|
||||
@@ -309,6 +373,12 @@ class TrainConfig:
|
||||
# unmasked reign. It is unmasked regularization basically
|
||||
self.inverted_mask_prior = kwargs.get('inverted_mask_prior', False)
|
||||
self.inverted_mask_prior_multiplier = kwargs.get('inverted_mask_prior_multiplier', 0.5)
|
||||
|
||||
# DOP will will run the same image and prompt through the network without the trigger word blank and use it as a target
|
||||
self.diff_output_preservation = kwargs.get('diff_output_preservation', False)
|
||||
self.diff_output_preservation_multiplier = kwargs.get('diff_output_preservation_multiplier', 1.0)
|
||||
# If the trigger word is in the prompt, we will use this class name to replace it eg. "sks woman" -> "woman"
|
||||
self.diff_output_preservation_class = kwargs.get('diff_output_preservation_class', '')
|
||||
|
||||
# legacy
|
||||
if match_adapter_assist and self.match_adapter_chance == 0.0:
|
||||
@@ -318,8 +388,8 @@ class TrainConfig:
|
||||
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")
|
||||
# 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)
|
||||
@@ -335,7 +405,7 @@ class TrainConfig:
|
||||
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')
|
||||
self.loss_type = kwargs.get('loss_type', 'mse') # mse, mae, wavelet
|
||||
|
||||
# scale the prediction by this. Increase for more detail, decrease for less
|
||||
self.pred_scaler = kwargs.get('pred_scaler', 1.0)
|
||||
@@ -347,7 +417,8 @@ class TrainConfig:
|
||||
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:
|
||||
# if it is set explicitly to false, leave it false.
|
||||
if ema_config is not None and ema_config.get('use_ema', None) is not None:
|
||||
ema_config['use_ema'] = True
|
||||
print(f"Using EMA")
|
||||
else:
|
||||
@@ -358,13 +429,40 @@ class TrainConfig:
|
||||
# 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.timestep_type = kwargs.get('timestep_type', 'sigmoid') # sigmoid, linear, lognorm_blend
|
||||
self.linear_timesteps = kwargs.get('linear_timesteps', False)
|
||||
self.linear_timesteps2 = kwargs.get('linear_timesteps2', False)
|
||||
self.disable_sampling = kwargs.get('disable_sampling', False)
|
||||
|
||||
# will cache a blank prompt or the trigger word, and unload the text encoder to cpu
|
||||
# will make training faster and use less vram
|
||||
self.unload_text_encoder = kwargs.get('unload_text_encoder', False)
|
||||
# for swapping which parameters are trained during training
|
||||
self.do_paramiter_swapping = kwargs.get('do_paramiter_swapping', False)
|
||||
# 0.1 is 10% of the parameters active at a time lower is less vram, higher is more
|
||||
self.paramiter_swapping_factor = kwargs.get('paramiter_swapping_factor', 0.1)
|
||||
# bypass the guidance embedding for training. For open flux with guidance embedding
|
||||
self.bypass_guidance_embedding = kwargs.get('bypass_guidance_embedding', False)
|
||||
|
||||
# diffusion feature extractor
|
||||
self.diffusion_feature_extractor_path = kwargs.get('diffusion_feature_extractor_path', None)
|
||||
self.diffusion_feature_extractor_weight = kwargs.get('diffusion_feature_extractor_weight', 1.0)
|
||||
|
||||
# optimal noise pairing
|
||||
self.optimal_noise_pairing_samples = kwargs.get('optimal_noise_pairing_samples', 1)
|
||||
|
||||
# forces same noise for the same image at a given size.
|
||||
self.force_consistent_noise = kwargs.get('force_consistent_noise', False)
|
||||
|
||||
|
||||
ModelArch = Literal['sd1', 'sd2', 'sd3', 'sdxl', 'pixart', 'pixart_sigma', 'auraflow', 'flux', 'flex2', 'lumina2', 'vega', 'ssd', 'wan21']
|
||||
|
||||
|
||||
class ModelConfig:
|
||||
def __init__(self, **kwargs):
|
||||
self.name_or_path: str = kwargs.get('name_or_path', None)
|
||||
# name or path is updated on fine tuning. Keep a copy of the original
|
||||
self.name_or_path_original: str = self.name_or_path
|
||||
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)
|
||||
@@ -372,6 +470,10 @@ class ModelConfig:
|
||||
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)
|
||||
self.is_flex2: bool = kwargs.get('is_flex2', False)
|
||||
if self.is_flex2:
|
||||
self.is_flux = True
|
||||
self.is_lumina2: bool = kwargs.get('is_lumina2', False)
|
||||
if self.is_pixart_sigma:
|
||||
self.is_pixart = True
|
||||
self.use_flux_cfg = kwargs.get('use_flux_cfg', False)
|
||||
@@ -386,6 +488,7 @@ class ModelConfig:
|
||||
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.inference_lora_path = kwargs.get('inference_lora_path', None)
|
||||
self.latent_space_version = kwargs.get('latent_space_version', None)
|
||||
|
||||
# only for SDXL models for now
|
||||
@@ -415,8 +518,81 @@ class ModelConfig:
|
||||
|
||||
# only for flux for now
|
||||
self.quantize = kwargs.get("quantize", False)
|
||||
self.quantize_te = kwargs.get("quantize_te", self.quantize)
|
||||
self.qtype = kwargs.get("qtype", "qfloat8")
|
||||
self.qtype_te = kwargs.get("qtype_te", "qfloat8")
|
||||
self.low_vram = kwargs.get("low_vram", False)
|
||||
pass
|
||||
self.attn_masking = kwargs.get("attn_masking", False)
|
||||
if self.attn_masking and not self.is_flux:
|
||||
raise ValueError("attn_masking is only supported with flux models currently")
|
||||
# for targeting a specific layers
|
||||
self.ignore_if_contains: Optional[List[str]] = kwargs.get("ignore_if_contains", None)
|
||||
self.only_if_contains: Optional[List[str]] = kwargs.get("only_if_contains", None)
|
||||
self.quantize_kwargs = kwargs.get("quantize_kwargs", {})
|
||||
|
||||
# splits the model over the available gpus WIP
|
||||
self.split_model_over_gpus = kwargs.get("split_model_over_gpus", False)
|
||||
if self.split_model_over_gpus and not self.is_flux:
|
||||
raise ValueError("split_model_over_gpus is only supported with flux models currently")
|
||||
self.split_model_other_module_param_count_scale = kwargs.get("split_model_other_module_param_count_scale", 0.3)
|
||||
|
||||
self.te_name_or_path = kwargs.get("te_name_or_path", None)
|
||||
|
||||
self.arch: ModelArch = kwargs.get("arch", None)
|
||||
|
||||
# handle migrating to new model arch
|
||||
if self.arch is not None:
|
||||
# reverse the arch to the old style
|
||||
if self.arch == 'sd2':
|
||||
self.is_v2 = True
|
||||
elif self.arch == 'sd3':
|
||||
self.is_v3 = True
|
||||
elif self.arch == 'sdxl':
|
||||
self.is_xl = True
|
||||
elif self.arch == 'pixart':
|
||||
self.is_pixart = True
|
||||
elif self.arch == 'pixart_sigma':
|
||||
self.is_pixart_sigma = True
|
||||
elif self.arch == 'auraflow':
|
||||
self.is_auraflow = True
|
||||
elif self.arch == 'flux':
|
||||
self.is_flux = True
|
||||
elif self.arch == 'flex2':
|
||||
self.is_flex2 = True
|
||||
elif self.arch == 'lumina2':
|
||||
self.is_lumina2 = True
|
||||
elif self.arch == 'vega':
|
||||
self.is_vega = True
|
||||
elif self.arch == 'ssd':
|
||||
self.is_ssd = True
|
||||
else:
|
||||
pass
|
||||
if self.arch is None:
|
||||
if kwargs.get('is_v2', False):
|
||||
self.arch = 'sd2'
|
||||
elif kwargs.get('is_v3', False):
|
||||
self.arch = 'sd3'
|
||||
elif kwargs.get('is_xl', False):
|
||||
self.arch = 'sdxl'
|
||||
elif kwargs.get('is_pixart', False):
|
||||
self.arch = 'pixart'
|
||||
elif kwargs.get('is_pixart_sigma', False):
|
||||
self.arch = 'pixart_sigma'
|
||||
elif kwargs.get('is_auraflow', False):
|
||||
self.arch = 'auraflow'
|
||||
elif kwargs.get('is_flux', False):
|
||||
self.arch = 'flux'
|
||||
elif kwargs.get('is_flex2', False):
|
||||
self.arch = 'flex2'
|
||||
elif kwargs.get('is_lumina2', False):
|
||||
self.arch = 'lumina2'
|
||||
elif kwargs.get('is_vega', False):
|
||||
self.arch = 'vega'
|
||||
elif kwargs.get('is_ssd', False):
|
||||
self.arch = 'ssd'
|
||||
else:
|
||||
self.arch = 'sd1'
|
||||
|
||||
|
||||
|
||||
class EMAConfig:
|
||||
@@ -425,6 +601,11 @@ class EMAConfig:
|
||||
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)
|
||||
|
||||
# every update, the params are multiplied by this amount
|
||||
# only use for things without a bias like lora
|
||||
# similar to a decay in an optimizer but the opposite
|
||||
self.param_multiplier: float = kwargs.get('param_multiplier', 1.0)
|
||||
|
||||
|
||||
class ReferenceDatasetConfig:
|
||||
@@ -513,6 +694,8 @@ class DatasetConfig:
|
||||
self.dataset_path: str = kwargs.get('dataset_path', None)
|
||||
|
||||
self.default_caption: str = kwargs.get('default_caption', None)
|
||||
# trigger word for just this dataset
|
||||
self.trigger_word: str = kwargs.get('trigger_word', None)
|
||||
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):
|
||||
@@ -538,7 +721,10 @@ class DatasetConfig:
|
||||
self.flip_x: bool = kwargs.get('flip_x', False)
|
||||
self.flip_y: bool = kwargs.get('flip_y', False)
|
||||
self.augments: List[str] = kwargs.get('augments', [])
|
||||
self.control_path: str = kwargs.get('control_path', None) # depth maps, etc
|
||||
self.control_path: Union[str,List[str]] = kwargs.get('control_path', None) # depth maps, etc
|
||||
# inpaint images should be webp/png images with alpha channel. The alpha 0 (invisible) section will
|
||||
# be the part conditioned to be inpainted. The alpha 1 (visible) section will be the part that is ignored
|
||||
self.inpaint_path: Union[str,List[str]] = kwargs.get('inpaint_path', None)
|
||||
# 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
|
||||
@@ -580,6 +766,8 @@ class DatasetConfig:
|
||||
|
||||
# ip adapter / reference dataset
|
||||
self.clip_image_path: str = kwargs.get('clip_image_path', None) # depth maps, etc
|
||||
# get the clip image randomly from the same folder as the image. Useful for folder grouped pairs.
|
||||
self.clip_image_from_same_folder: bool = kwargs.get('clip_image_from_same_folder', False)
|
||||
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', [])
|
||||
@@ -591,6 +779,22 @@ class DatasetConfig:
|
||||
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)
|
||||
|
||||
# for video
|
||||
# if num_frames is greater than 1, the dataloader will look for video files.
|
||||
# num_frames will be the number of frames in the training batch. If num_frames is 1, it will look for images
|
||||
self.num_frames: int = kwargs.get('num_frames', 1)
|
||||
# if true, will shrink video to our frames. For instance, if we have a video with 100 frames and num_frames is 10,
|
||||
# we would pull frame 0, 10, 20, 30, 40, 50, 60, 70, 80, 90 so they are evenly spaced
|
||||
self.shrink_video_to_frames: bool = kwargs.get('shrink_video_to_frames', True)
|
||||
# fps is only used if shrink_video_to_frames is false. This will attempt to pull the num_frames at the given fps
|
||||
# it will select a random start frame and pull the frames at the given fps
|
||||
# this could have various issues with shorter videos and videos with variable fps
|
||||
# I recommend trimming your videos to the desired length and using shrink_video_to_frames(default)
|
||||
self.fps: int = kwargs.get('fps', 16)
|
||||
|
||||
# debug the frame count and frame selection. You dont need this. It is for debugging.
|
||||
self.debug: bool = kwargs.get('debug', False)
|
||||
|
||||
|
||||
def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]:
|
||||
@@ -640,6 +844,10 @@ class GenerateImageConfig:
|
||||
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
|
||||
logger: Optional[EmptyLogger] = None,
|
||||
num_frames: int = 1,
|
||||
fps: int = 15,
|
||||
ctrl_idx: int = 0
|
||||
):
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
@@ -668,6 +876,10 @@ class GenerateImageConfig:
|
||||
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 []
|
||||
self.num_frames = num_frames
|
||||
self.fps = fps
|
||||
self.ctrl_idx = ctrl_idx
|
||||
|
||||
|
||||
# prompt string will override any settings above
|
||||
self._process_prompt_string()
|
||||
@@ -697,6 +909,8 @@ class GenerateImageConfig:
|
||||
self.height = max(64, self.height - self.height % 8) # round to divisible by 8
|
||||
self.width = max(64, self.width - self.width % 8) # round to divisible by 8
|
||||
|
||||
self.logger = logger
|
||||
|
||||
def set_gen_time(self, gen_time: int = None):
|
||||
if gen_time is not None:
|
||||
self.gen_time = gen_time
|
||||
@@ -732,11 +946,30 @@ class GenerateImageConfig:
|
||||
# make parent dirs
|
||||
os.makedirs(self.output_folder, exist_ok=True)
|
||||
self.set_gen_time()
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
if isinstance(image, list):
|
||||
# video
|
||||
if self.num_frames == 1:
|
||||
raise ValueError(f"Expected 1 img but got a list {len(image)}")
|
||||
if self.output_ext == 'webp':
|
||||
# save as animated webp
|
||||
duration = 1000 // self.fps # Convert fps to milliseconds per frame
|
||||
image[0].save(
|
||||
self.get_image_path(count, max_count),
|
||||
format='WEBP',
|
||||
append_images=image[1:],
|
||||
save_all=True,
|
||||
duration=duration, # Duration per frame in milliseconds
|
||||
loop=0, # 0 means loop forever
|
||||
quality=80 # Quality setting (0-100)
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported video format {self.output_ext}")
|
||||
else:
|
||||
# TODO save image gen header info for A1111 and us, our seeds probably wont match
|
||||
image.save(self.get_image_path(count, max_count))
|
||||
# do prompt file
|
||||
if self.add_prompt_file:
|
||||
self.save_prompt_file(count, max_count)
|
||||
|
||||
def save_prompt_file(self, count: int = 0, max_count=0):
|
||||
# save prompt file
|
||||
@@ -757,7 +990,10 @@ class GenerateImageConfig:
|
||||
prompt += ' --gr ' + str(self.guidance_rescale)
|
||||
|
||||
# get gen info
|
||||
f.write(self.prompt)
|
||||
try:
|
||||
f.write(self.prompt)
|
||||
except Exception as e:
|
||||
print(f"Error writing prompt file. Prompt contains non-unicode characters. {e}")
|
||||
|
||||
def _process_prompt_string(self):
|
||||
# we will try to support all sd-scripts where we can
|
||||
@@ -832,6 +1068,12 @@ class GenerateImageConfig:
|
||||
elif flag == 'extra_values':
|
||||
# split by comma
|
||||
self.extra_values = [float(val) for val in content.split(',')]
|
||||
elif flag == 'frames':
|
||||
self.num_frames = int(content)
|
||||
elif flag == 'fps':
|
||||
self.fps = int(content)
|
||||
elif flag == 'ctrl_idx':
|
||||
self.ctrl_idx = int(content)
|
||||
|
||||
def post_process_embeddings(
|
||||
self,
|
||||
@@ -840,3 +1082,23 @@ class GenerateImageConfig:
|
||||
):
|
||||
# this is called after prompt embeds are encoded. We can override them in the future here
|
||||
pass
|
||||
|
||||
def log_image(self, image, count: int = 0, max_count=0):
|
||||
if self.logger is None:
|
||||
return
|
||||
|
||||
self.logger.log_image(image, count, self.prompt)
|
||||
|
||||
|
||||
def validate_configs(
|
||||
train_config: TrainConfig,
|
||||
model_config: ModelConfig,
|
||||
save_config: SaveConfig,
|
||||
):
|
||||
if model_config.is_flux:
|
||||
if save_config.save_format != 'diffusers':
|
||||
# make it diffusers
|
||||
save_config.save_format = 'diffusers'
|
||||
if model_config.use_flux_cfg:
|
||||
# bypass the embedding
|
||||
train_config.bypass_guidance_embedding = True
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import math
|
||||
import torch
|
||||
import sys
|
||||
|
||||
@@ -6,17 +7,24 @@ from torch.nn import Parameter
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection, T5EncoderModel, CLIPTextModel, \
|
||||
CLIPTokenizer, T5Tokenizer
|
||||
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.models.clip_fusion import CLIPFusionModule
|
||||
from toolkit.models.clip_pre_processor import CLIPImagePreProcessor
|
||||
from toolkit.models.control_lora_adapter import ControlLoraAdapter
|
||||
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.models.redux import ReduxImageEncoder
|
||||
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
|
||||
from toolkit.models.pixtral_vision import PixtralVisionEncoderCompatible, PixtralVisionImagePreprocessorCompatible
|
||||
import random
|
||||
|
||||
from toolkit.util.mask import generate_random_mask
|
||||
|
||||
sys.path.append(REPOS_ROOT)
|
||||
from typing import TYPE_CHECKING, Union, Iterator, Mapping, Any, Tuple, List, Optional, Dict
|
||||
@@ -25,7 +33,7 @@ from ipadapter.ip_adapter.attention_processor import AttnProcessor, IPAttnProces
|
||||
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.config_modules import AdapterConfig, AdapterTypes, TrainConfig
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
import weakref
|
||||
|
||||
@@ -40,7 +48,7 @@ from transformers import (
|
||||
ConvNextModel,
|
||||
ConvNextForImageClassification,
|
||||
ConvNextImageProcessor,
|
||||
UMT5EncoderModel, LlamaTokenizerFast
|
||||
UMT5EncoderModel, LlamaTokenizerFast, AutoModel, AutoTokenizer, BitsAndBytesConfig
|
||||
)
|
||||
from toolkit.models.size_agnostic_feature_encoder import SAFEImageProcessor, SAFEVisionModel
|
||||
|
||||
@@ -48,14 +56,17 @@ from transformers import ViTHybridImageProcessor, ViTHybridForImageClassificatio
|
||||
|
||||
from transformers import ViTFeatureExtractor, ViTForImageClassification
|
||||
|
||||
from toolkit.models.llm_adapter import LLMAdapter
|
||||
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class CustomAdapter(torch.nn.Module):
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: 'AdapterConfig'):
|
||||
def __init__(self, sd: 'StableDiffusion', adapter_config: 'AdapterConfig', train_config: 'TrainConfig'):
|
||||
super().__init__()
|
||||
self.config = adapter_config
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.train_config = train_config
|
||||
self.device = self.sd_ref().unet.device
|
||||
self.image_processor: CLIPImageProcessor = None
|
||||
self.input_size = 224
|
||||
@@ -73,7 +84,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
|
||||
self.position_ids: Optional[List[int]] = None
|
||||
|
||||
self.num_control_images = 1
|
||||
self.num_control_images = self.config.num_control_images
|
||||
self.token_mask: Optional[torch.Tensor] = None
|
||||
|
||||
# setup clip
|
||||
@@ -90,6 +101,9 @@ class CustomAdapter(torch.nn.Module):
|
||||
self.te_augmenter: TEAugAdapter = None
|
||||
self.vd_adapter: VisionDirectAdapter = None
|
||||
self.single_value_adapter: SingleValueAdapter = None
|
||||
self.redux_adapter: ReduxImageEncoder = None
|
||||
self.control_lora: ControlLoraAdapter = None
|
||||
|
||||
self.conditional_embeds: Optional[torch.Tensor] = None
|
||||
self.unconditional_embeds: Optional[torch.Tensor] = None
|
||||
|
||||
@@ -117,11 +131,11 @@ class CustomAdapter(torch.nn.Module):
|
||||
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']
|
||||
embed_dim = sd.unet_unwrapped.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']
|
||||
embed_dim = sd.unet_unwrapped.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':
|
||||
@@ -192,12 +206,53 @@ class CustomAdapter(torch.nn.Module):
|
||||
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 == 'llm_adapter':
|
||||
kwargs = {}
|
||||
if self.config.quantize_llm:
|
||||
bnb_kwargs = {
|
||||
'load_in_4bit': True,
|
||||
'bnb_4bit_quant_type': "nf4",
|
||||
'bnb_4bit_compute_dtype': torch.bfloat16
|
||||
}
|
||||
quantization_config = BitsAndBytesConfig(**bnb_kwargs)
|
||||
kwargs['quantization_config'] = quantization_config
|
||||
kwargs['torch_dtype'] = torch_dtype
|
||||
self.te = AutoModel.from_pretrained(
|
||||
self.config.text_encoder_path,
|
||||
**kwargs
|
||||
)
|
||||
else:
|
||||
self.te = AutoModel.from_pretrained(self.config.text_encoder_path).to(
|
||||
self.sd_ref().unet.device,
|
||||
dtype=torch_dtype,
|
||||
)
|
||||
self.te.to = lambda *args, **kwargs: None
|
||||
self.te.eval()
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(self.config.text_encoder_path)
|
||||
self.llm_adapter = LLMAdapter(
|
||||
adapter=self,
|
||||
sd=self.sd_ref(),
|
||||
llm=self.te,
|
||||
tokenizer=self.tokenizer,
|
||||
num_cloned_blocks=self.config.num_cloned_blocks,
|
||||
)
|
||||
self.llm_adapter.to(self.device, torch_dtype)
|
||||
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)
|
||||
elif self.adapter_type == 'redux':
|
||||
vision_hidden_size = self.vision_encoder.config.hidden_size
|
||||
self.redux_adapter = ReduxImageEncoder(vision_hidden_size, 4096, self.device, torch_dtype)
|
||||
elif self.adapter_type == 'control_lora':
|
||||
self.control_lora = ControlLoraAdapter(
|
||||
self,
|
||||
sd=self.sd_ref(),
|
||||
config=self.config,
|
||||
train_config=self.train_config
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown adapter type: {self.adapter_type}")
|
||||
|
||||
@@ -229,7 +284,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
def setup_clip(self):
|
||||
adapter_config = self.config
|
||||
sd = self.sd_ref()
|
||||
if self.config.type == "text_encoder" or self.config.type == "single_value":
|
||||
if self.config.type in ["text_encoder", "llm_adapter", "single_value", "control_lora"]:
|
||||
return
|
||||
if self.config.type == 'photo_maker':
|
||||
try:
|
||||
@@ -257,6 +312,22 @@ class CustomAdapter(torch.nn.Module):
|
||||
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 == 'siglip2':
|
||||
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 == 'pixtral':
|
||||
self.image_processor = PixtralVisionImagePreprocessorCompatible(
|
||||
max_image_size=self.config.pixtral_max_image_size,
|
||||
)
|
||||
self.vision_encoder = PixtralVisionEncoderCompatible.from_pretrained(
|
||||
adapter_config.image_encoder_path,
|
||||
).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)
|
||||
@@ -272,7 +343,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
self.vision_encoder = SAFEVisionModel(
|
||||
in_channels=3,
|
||||
num_tokens=self.config.safe_tokens,
|
||||
num_vectors=sd.unet.config['cross_attention_dim'],
|
||||
num_vectors=sd.unet_unwrapped.config['cross_attention_dim'],
|
||||
reducer_channels=self.config.safe_reducer_channels,
|
||||
channels=self.config.safe_channels,
|
||||
downscale_factor=8
|
||||
@@ -390,6 +461,9 @@ class CustomAdapter(torch.nn.Module):
|
||||
|
||||
if 'te_adapter' in state_dict:
|
||||
self.te_adapter.load_state_dict(state_dict['te_adapter'], strict=strict)
|
||||
|
||||
if 'llm_adapter' in state_dict:
|
||||
self.llm_adapter.load_state_dict(state_dict['llm_adapter'], strict=strict)
|
||||
|
||||
if 'te_augmenter' in state_dict:
|
||||
self.te_augmenter.load_state_dict(state_dict['te_augmenter'], strict=strict)
|
||||
@@ -397,12 +471,12 @@ class CustomAdapter(torch.nn.Module):
|
||||
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)
|
||||
self.vd_adapter.load_state_dict(state_dict['dvadapter'], strict=False)
|
||||
|
||||
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:
|
||||
if 'vision_encoder' in state_dict:
|
||||
self.vision_encoder.load_state_dict(state_dict['vision_encoder'], strict=strict)
|
||||
|
||||
if 'fuse_module' in state_dict:
|
||||
@@ -413,6 +487,21 @@ class CustomAdapter(torch.nn.Module):
|
||||
self.ilora_module.load_state_dict(state_dict['ilora'], strict=strict)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
if 'redux_up' in state_dict:
|
||||
# state dict is seperated. so recombine it
|
||||
new_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
for k2, v2 in v.items():
|
||||
new_dict[k + '.' + k2] = v2
|
||||
self.redux_adapter.load_state_dict(new_dict, strict=True)
|
||||
|
||||
if self.adapter_type == 'control_lora':
|
||||
# state dict is seperated. so recombine it
|
||||
new_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
for k2, v2 in v.items():
|
||||
new_dict[k + '.' + k2] = v2
|
||||
self.control_lora.load_weights(new_dict, strict=strict)
|
||||
|
||||
pass
|
||||
|
||||
@@ -438,6 +527,9 @@ class CustomAdapter(torch.nn.Module):
|
||||
elif self.adapter_type == 'text_encoder':
|
||||
state_dict["te_adapter"] = self.te_adapter.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'llm_adapter':
|
||||
state_dict["llm_adapter"] = self.llm_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()
|
||||
@@ -445,8 +537,8 @@ class CustomAdapter(torch.nn.Module):
|
||||
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()
|
||||
# if self.config.train_image_encoder: # always return vision 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()
|
||||
@@ -456,6 +548,16 @@ class CustomAdapter(torch.nn.Module):
|
||||
state_dict["vision_encoder"] = self.vision_encoder.state_dict()
|
||||
state_dict["ilora"] = self.ilora_module.state_dict()
|
||||
return state_dict
|
||||
elif self.adapter_type == 'redux':
|
||||
d = self.redux_adapter.state_dict()
|
||||
for k, v in d.items():
|
||||
state_dict[k] = v
|
||||
return state_dict
|
||||
elif self.adapter_type == 'control_lora':
|
||||
d = self.control_lora.get_state_dict()
|
||||
for k, v in d.items():
|
||||
state_dict[k] = v
|
||||
return state_dict
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -465,6 +567,135 @@ class CustomAdapter(torch.nn.Module):
|
||||
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_noisy_latents(self, latents: torch.Tensor, batch:DataLoaderBatchDTO):
|
||||
with torch.no_grad():
|
||||
if self.adapter_type in ['control_lora']:
|
||||
# inpainting input is 0-1 (bs, 4, h, w) on batch.inpaint_tensor
|
||||
# 4th channel is the mask with 1 being keep area and 0 being area to inpaint.
|
||||
sd: StableDiffusion = self.sd_ref()
|
||||
inpainting_latent = None
|
||||
if self.config.has_inpainting_input:
|
||||
do_dropout = random.random() < self.config.control_image_dropout
|
||||
# do random mask if we dont have one
|
||||
inpaint_tensor = batch.inpaint_tensor
|
||||
if inpaint_tensor is None and not do_dropout:
|
||||
# generate a random one since we dont have one
|
||||
# this will make random blobs, invert the blobs for now as we normanlly inpaint the alpha
|
||||
inpaint_tensor = 1 - generate_random_mask(
|
||||
batch_size=latents.shape[0],
|
||||
height=latents.shape[2],
|
||||
width=latents.shape[3],
|
||||
device=latents.device,
|
||||
).to(latents.device, latents.dtype)
|
||||
if inpaint_tensor is not None and not do_dropout:
|
||||
|
||||
if inpaint_tensor.shape[1] == 4:
|
||||
# get just the mask
|
||||
inpainting_tensor_mask = inpaint_tensor[:, 3:4, :, :].to(latents.device, dtype=latents.dtype)
|
||||
elif inpaint_tensor.shape[1] == 3:
|
||||
# rgb mask. Just get one channel
|
||||
inpainting_tensor_mask = inpaint_tensor[:, 0:1, :, :].to(latents.device, dtype=latents.dtype)
|
||||
else:
|
||||
inpainting_tensor_mask = inpaint_tensor
|
||||
|
||||
# # use our batch latents so we cna avoid ancoding again
|
||||
inpainting_latent = batch.latents
|
||||
|
||||
# resize the mask to match the new encoded size
|
||||
inpainting_tensor_mask = F.interpolate(inpainting_tensor_mask, size=(inpainting_latent.shape[2], inpainting_latent.shape[3]), mode='bilinear')
|
||||
inpainting_tensor_mask = inpainting_tensor_mask.to(latents.device, latents.dtype)
|
||||
|
||||
do_mask_invert = False
|
||||
if self.config.invert_inpaint_mask_chance > 0.0:
|
||||
do_mask_invert = random.random() < self.config.invert_inpaint_mask_chance
|
||||
if do_mask_invert:
|
||||
# invert the mask
|
||||
inpainting_tensor_mask = 1 - inpainting_tensor_mask
|
||||
|
||||
# mask out the inpainting area, it is currently 0 for inpaint area, and 1 for keep area
|
||||
# we are zeroing our the latents in the inpaint area not on the pixel space.
|
||||
inpainting_latent = inpainting_latent * inpainting_tensor_mask
|
||||
|
||||
# mask needs to be 1 for inpaint area and 0 for area to leave alone. So flip it.
|
||||
inpainting_tensor_mask = 1 - inpainting_tensor_mask
|
||||
# leave the mask as 0-1 and concat on channel of latents
|
||||
inpainting_latent = torch.cat((inpainting_latent, inpainting_tensor_mask), dim=1)
|
||||
else:
|
||||
# we have iinpainting but didnt get a control. or we are doing a dropout
|
||||
# the input needs to be all zeros for the latents and all 1s for the mask
|
||||
inpainting_latent = torch.zeros_like(latents)
|
||||
# add ones for the mask since we are technically inpainting everything
|
||||
inpainting_latent = torch.cat((inpainting_latent, torch.ones_like(inpainting_latent[:, :1, :, :])), dim=1)
|
||||
|
||||
if self.config.num_control_images == 1:
|
||||
# this is our only control
|
||||
control_latent = inpainting_latent.to(latents.device, latents.dtype)
|
||||
latents = torch.cat((latents, control_latent), dim=1)
|
||||
return latents.detach()
|
||||
|
||||
control_tensor = batch.control_tensor.to(latents.device, dtype=latents.dtype)
|
||||
if control_tensor is None:
|
||||
# concat random normal noise onto the latents
|
||||
# check dimension, this is before they are rearranged
|
||||
# it is latent_model_input = torch.cat([latents, control_image], dim=2) after rearranging
|
||||
ctrl = torch.zeros(
|
||||
latents.shape[0], # bs
|
||||
latents.shape[1] * self.num_control_images, # ch
|
||||
latents.shape[2],
|
||||
latents.shape[3],
|
||||
device=latents.device,
|
||||
dtype=latents.dtype
|
||||
)
|
||||
if inpainting_latent is not None:
|
||||
# inpainting always comes first
|
||||
ctrl = torch.cat((inpainting_latent, ctrl), dim=1)
|
||||
latents = torch.cat((latents, ctrl), dim=1)
|
||||
return latents.detach()
|
||||
# if we have multiple control tensors, they come in like [bs, num_control_images, ch, h, w]
|
||||
# if we have 1, it comes in like [bs, ch, h, w]
|
||||
# stack out control tensors to be [bs, ch * num_control_images, h, w]
|
||||
|
||||
control_tensor_list = []
|
||||
if len(control_tensor.shape) == 4:
|
||||
control_tensor_list.append(control_tensor)
|
||||
else:
|
||||
# reshape
|
||||
control_tensor = control_tensor.view(
|
||||
control_tensor.shape[0],
|
||||
control_tensor.shape[1] * control_tensor.shape[2],
|
||||
control_tensor.shape[3],
|
||||
control_tensor.shape[4]
|
||||
)
|
||||
control_tensor_list = control_tensor.chunk(self.num_control_images, dim=1)
|
||||
control_latent_list = []
|
||||
for control_tensor in control_tensor_list:
|
||||
do_dropout = random.random() < self.config.control_image_dropout
|
||||
if do_dropout:
|
||||
# dropout with noise
|
||||
control_latent_list.append(torch.zeros_like(batch.latents))
|
||||
else:
|
||||
# it is 0-1 need to convert to -1 to 1
|
||||
control_tensor = control_tensor * 2 - 1
|
||||
|
||||
control_tensor = control_tensor.to(sd.vae_device_torch, dtype=sd.torch_dtype)
|
||||
|
||||
# if it is not the size of batch.tensor, (bs,ch,h,w) then we need to resize it
|
||||
if control_tensor.shape[2] != batch.tensor.shape[2] or control_tensor.shape[3] != batch.tensor.shape[3]:
|
||||
control_tensor = F.interpolate(control_tensor, size=(batch.tensor.shape[2], batch.tensor.shape[3]), mode='bicubic')
|
||||
|
||||
# encode it
|
||||
control_latent = sd.encode_images(control_tensor).to(latents.device, latents.dtype)
|
||||
control_latent_list.append(control_latent)
|
||||
# stack them on the channel dimension
|
||||
control_latent = torch.cat(control_latent_list, dim=1)
|
||||
if inpainting_latent is not None:
|
||||
# inpainting always comes first
|
||||
control_latent = torch.cat((inpainting_latent, control_latent), dim=1)
|
||||
# concat it onto the latents
|
||||
latents = torch.cat((latents, control_latent), dim=1)
|
||||
return latents.detach()
|
||||
return latents
|
||||
|
||||
|
||||
def condition_prompt(
|
||||
@@ -472,7 +703,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
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':
|
||||
if self.adapter_type in ['clip_fusion', 'ilora', 'vision_direct', 'redux', 'control_lora']:
|
||||
return prompt
|
||||
elif self.adapter_type == 'text_encoder':
|
||||
# todo allow for training
|
||||
@@ -482,6 +713,14 @@ class CustomAdapter(torch.nn.Module):
|
||||
self.unconditional_embeds = self.te_adapter.encode_text(prompt).detach()
|
||||
else:
|
||||
self.conditional_embeds = self.te_adapter.encode_text(prompt).detach()
|
||||
elif self.adapter_type == 'llm_adapter':
|
||||
# todo allow for training
|
||||
with torch.no_grad():
|
||||
# encode and save the embeds
|
||||
if is_unconditional:
|
||||
self.unconditional_embeds = self.llm_adapter.encode_text(prompt).detach()
|
||||
else:
|
||||
self.conditional_embeds = self.llm_adapter.encode_text(prompt).detach()
|
||||
return prompt
|
||||
elif self.adapter_type == 'photo_maker':
|
||||
if is_unconditional:
|
||||
@@ -585,16 +824,25 @@ class CustomAdapter(torch.nn.Module):
|
||||
quad_count=4,
|
||||
is_generating_samples=False,
|
||||
) -> PromptEmbeds:
|
||||
if self.adapter_type == 'text_encoder' and is_generating_samples:
|
||||
if self.adapter_type == 'text_encoder':
|
||||
# replace the prompt embed with ours
|
||||
if is_unconditional:
|
||||
return self.unconditional_embeds.clone()
|
||||
return self.conditional_embeds.clone()
|
||||
if self.adapter_type == 'llm_adapter':
|
||||
# replace the prompt embed with ours
|
||||
if is_unconditional:
|
||||
prompt_embeds.text_embeds = self.unconditional_embeds.text_embeds.clone()
|
||||
prompt_embeds.attention_mask = self.unconditional_embeds.attention_mask.clone()
|
||||
return prompt_embeds
|
||||
prompt_embeds.text_embeds = self.conditional_embeds.text_embeds.clone()
|
||||
prompt_embeds.attention_mask = self.conditional_embeds.attention_mask.clone()
|
||||
return prompt_embeds
|
||||
|
||||
if self.adapter_type == 'ilora':
|
||||
return prompt_embeds
|
||||
|
||||
if self.adapter_type == 'photo_maker' or self.adapter_type == 'clip_fusion':
|
||||
if self.adapter_type == 'photo_maker' or self.adapter_type == 'clip_fusion' or self.adapter_type == 'redux':
|
||||
if is_unconditional:
|
||||
# we dont condition the negative embeds for photo maker
|
||||
return prompt_embeds.clone()
|
||||
@@ -616,6 +864,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
do_convert_rgb=True
|
||||
).pixel_values
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
@@ -696,13 +945,49 @@ class CustomAdapter(torch.nn.Module):
|
||||
)
|
||||
return prompt_embeds
|
||||
|
||||
elif self.adapter_type == 'redux':
|
||||
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()
|
||||
|
||||
img_embeds = self.redux_adapter(img_embeds.to(self.device, get_torch_dtype(self.sd_ref().dtype)))
|
||||
|
||||
prompt_embeds.text_embeds = torch.cat((prompt_embeds.text_embeds, img_embeds), dim=-2)
|
||||
return prompt_embeds
|
||||
else:
|
||||
return prompt_embeds
|
||||
|
||||
def get_empty_clip_image(self, batch_size: int) -> torch.Tensor:
|
||||
def get_empty_clip_image(self, batch_size: int, shape=None) -> torch.Tensor:
|
||||
with torch.no_grad():
|
||||
tensors_0_1 = torch.rand([batch_size, 3, self.input_size, self.input_size], device=self.device)
|
||||
if shape is None:
|
||||
shape = [batch_size, 3, self.input_size, self.input_size]
|
||||
tensors_0_1 = torch.rand(shape, 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
|
||||
@@ -720,8 +1005,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
def train(self, mode: bool = True):
|
||||
if self.config.train_image_encoder:
|
||||
self.vision_encoder.train(mode)
|
||||
else:
|
||||
super().train(mode)
|
||||
super().train(mode)
|
||||
|
||||
def trigger_pre_te(
|
||||
self,
|
||||
@@ -732,6 +1016,7 @@ class CustomAdapter(torch.nn.Module):
|
||||
batch_size=1,
|
||||
) -> PromptEmbeds:
|
||||
if self.adapter_type == 'ilora' or self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter':
|
||||
skip_unconditional = self.sd_ref().is_flux
|
||||
if tensors_0_1 is None:
|
||||
tensors_0_1 = self.get_empty_clip_image(batch_size)
|
||||
has_been_preprocessed = True
|
||||
@@ -757,11 +1042,37 @@ class CustomAdapter(torch.nn.Module):
|
||||
).pixel_values
|
||||
else:
|
||||
clip_image = tensors_0_1
|
||||
|
||||
# if is pixtral
|
||||
if self.config.image_encoder_arch == 'pixtral' and self.config.pixtral_random_image_size:
|
||||
# get the random size
|
||||
random_size = random.randint(256, self.config.pixtral_max_image_size)
|
||||
# images are already sized for max size, we have to fit them to the pixtral patch size to reduce / enlarge it farther.
|
||||
h, w = clip_image.shape[2], clip_image.shape[3]
|
||||
current_base_size = int(math.sqrt(w * h))
|
||||
ratio = current_base_size / random_size
|
||||
if ratio > 1:
|
||||
w = round(w / ratio)
|
||||
h = round(h / ratio)
|
||||
|
||||
width_tokens = (w - 1) // self.image_processor.image_patch_size + 1
|
||||
height_tokens = (h - 1) // self.image_processor.image_patch_size + 1
|
||||
assert width_tokens > 0
|
||||
assert height_tokens > 0
|
||||
|
||||
new_image_size = (
|
||||
width_tokens * self.image_processor.image_patch_size,
|
||||
height_tokens * self.image_processor.image_patch_size,
|
||||
)
|
||||
|
||||
# resize the image
|
||||
clip_image = F.interpolate(clip_image, size=new_image_size, mode='bicubic', align_corners=False)
|
||||
|
||||
|
||||
batch_size = clip_image.shape[0]
|
||||
if self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter':
|
||||
if (self.adapter_type == 'vision_direct' or self.adapter_type == 'te_augmenter') and not skip_unconditional:
|
||||
# add an unconditional so we can save it
|
||||
unconditional = self.get_empty_clip_image(batch_size).to(
|
||||
unconditional = self.get_empty_clip_image(batch_size, shape=clip_image.shape).to(
|
||||
clip_image.device, dtype=clip_image.dtype
|
||||
)
|
||||
clip_image = torch.cat([unconditional, clip_image], dim=0)
|
||||
@@ -840,11 +1151,14 @@ class CustomAdapter(torch.nn.Module):
|
||||
elif self.config.clip_layer == 'last_hidden_state':
|
||||
clip_image_embeds = clip_output.hidden_states[-1]
|
||||
else:
|
||||
clip_image_embeds = clip_output.image_embeds
|
||||
if hasattr(clip_output, 'image_embeds'):
|
||||
clip_image_embeds = clip_output.image_embeds
|
||||
elif hasattr(clip_output, 'pooler_output'):
|
||||
clip_image_embeds = clip_output.pooler_output
|
||||
# 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
|
||||
# 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()
|
||||
@@ -857,7 +1171,10 @@ class CustomAdapter(torch.nn.Module):
|
||||
|
||||
# save them to the conditional and unconditional
|
||||
try:
|
||||
self.unconditional_embeds, self.conditional_embeds = clip_image_embeds.chunk(2, dim=0)
|
||||
if skip_unconditional:
|
||||
self.unconditional_embeds, self.conditional_embeds = None, clip_image_embeds
|
||||
else:
|
||||
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}")
|
||||
|
||||
@@ -880,17 +1197,35 @@ class CustomAdapter(torch.nn.Module):
|
||||
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 == 'llm_adapter':
|
||||
yield from self.llm_adapter.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)
|
||||
if self.config.train_scaler:
|
||||
# only yield the self.block_scaler = torch.nn.Parameter(torch.tensor([1.0] * num_modules)
|
||||
yield self.vd_adapter.block_scaler
|
||||
else:
|
||||
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)
|
||||
if self.vd_adapter.resampler is not None:
|
||||
yield from self.vd_adapter.resampler.parameters(recurse)
|
||||
if self.vd_adapter.pool is not None:
|
||||
yield from self.vd_adapter.pool.parameters(recurse)
|
||||
if self.vd_adapter.sparse_autoencoder is not None:
|
||||
yield from self.vd_adapter.sparse_autoencoder.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)
|
||||
elif self.config.type == 'redux':
|
||||
yield from self.redux_adapter.parameters(recurse)
|
||||
elif self.config.type == 'control_lora':
|
||||
param_list = self.control_lora.get_params()
|
||||
for param in param_list:
|
||||
yield param
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -908,4 +1243,10 @@ class CustomAdapter(torch.nn.Module):
|
||||
additional[k] = v
|
||||
additional['clip_layer'] = self.config.clip_layer
|
||||
additional['image_encoder_arch'] = self.config.head_dim
|
||||
return additional
|
||||
return additional
|
||||
|
||||
def post_weight_update(self):
|
||||
# do any kind of updates after the weight update
|
||||
if self.config.type == 'vision_direct':
|
||||
self.vd_adapter.post_weight_update()
|
||||
pass
|
||||
@@ -20,6 +20,8 @@ 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, CLIPCachingMixin
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.accelerator import get_accelerator
|
||||
|
||||
import platform
|
||||
|
||||
@@ -28,6 +30,10 @@ def is_native_windows():
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
|
||||
image_extensions = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
video_extensions = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']
|
||||
|
||||
|
||||
class RescaleTransform:
|
||||
@@ -90,7 +96,7 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
|
||||
|
||||
# this might take a while
|
||||
print(f" - Preprocessing image dimensions")
|
||||
print_acc(f" - Preprocessing image dimensions")
|
||||
new_file_list = []
|
||||
bad_count = 0
|
||||
for file in tqdm(self.file_list):
|
||||
@@ -102,8 +108,8 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
|
||||
self.file_list = new_file_list
|
||||
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
print(f" - Found {bad_count} images that are too small")
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
print_acc(f" - Found {bad_count} images that are too small")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.path}"
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
@@ -128,8 +134,8 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
try:
|
||||
img = exif_transpose(Image.open(img_path)).convert('RGB')
|
||||
except Exception as e:
|
||||
print(f"Error opening image: {img_path}")
|
||||
print(e)
|
||||
print_acc(f"Error opening image: {img_path}")
|
||||
print_acc(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))
|
||||
|
||||
@@ -140,7 +146,7 @@ class ImageDataset(Dataset, CaptionMixin):
|
||||
if self.random_crop:
|
||||
if self.random_scale and min_img_size > self.resolution:
|
||||
if min_img_size < self.resolution:
|
||||
print(
|
||||
print_acc(
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={img_path}")
|
||||
scale_size = self.resolution
|
||||
else:
|
||||
@@ -243,11 +249,11 @@ class PairedImageDataset(Dataset):
|
||||
matched_files = [t for t in (set(tuple(i) for i in matched_files))]
|
||||
|
||||
self.file_list = matched_files
|
||||
print(f" - Found {len(self.file_list)} matching pairs")
|
||||
print_acc(f" - Found {len(self.file_list)} matching pairs")
|
||||
else:
|
||||
self.file_list = [os.path.join(self.path, file) for file in os.listdir(self.path) if
|
||||
file.lower().endswith(supported_exts)]
|
||||
print(f" - Found {len(self.file_list)} images")
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
|
||||
self.transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
@@ -374,8 +380,9 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
batch_size=1,
|
||||
sd: 'StableDiffusion' = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dataset_config = dataset_config
|
||||
self.is_video = dataset_config.num_frames > 1
|
||||
super().__init__()
|
||||
folder_path = dataset_config.folder_path
|
||||
self.dataset_path = dataset_config.dataset_path
|
||||
if self.dataset_path is None:
|
||||
@@ -405,7 +412,11 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
|
||||
# check if dataset_path is a folder or json
|
||||
if os.path.isdir(self.dataset_path):
|
||||
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'))]
|
||||
extensions = image_extensions
|
||||
if self.is_video:
|
||||
# only look for videos
|
||||
extensions = video_extensions
|
||||
file_list = [os.path.join(root, file) for root, _, files in os.walk(self.dataset_path) for file in files if file.lower().endswith(tuple(extensions))]
|
||||
else:
|
||||
# assume json
|
||||
with open(self.dataset_path, 'r') as f:
|
||||
@@ -435,17 +446,34 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
])
|
||||
|
||||
# this might take a while
|
||||
print(f"Dataset: {self.dataset_path}")
|
||||
print(f" - Preprocessing image dimensions")
|
||||
print_acc(f"Dataset: {self.dataset_path}")
|
||||
if self.is_video:
|
||||
print_acc(f" - Preprocessing video dimensions")
|
||||
else:
|
||||
print_acc(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')
|
||||
dataloader_version = "0.1.1"
|
||||
if os.path.exists(dataset_size_file):
|
||||
with open(dataset_size_file, 'r') as f:
|
||||
self.size_database = json.load(f)
|
||||
try:
|
||||
with open(dataset_size_file, 'r') as f:
|
||||
self.size_database = json.load(f)
|
||||
|
||||
if "__version__" not in self.size_database or self.size_database["__version__"] != dataloader_version:
|
||||
print_acc("Upgrading size database to new version")
|
||||
# old version, delete and recreate
|
||||
self.size_database = {}
|
||||
except Exception as e:
|
||||
print_acc(f"Error loading size database: {dataset_size_file}")
|
||||
print_acc(e)
|
||||
self.size_database = {}
|
||||
else:
|
||||
self.size_database = {}
|
||||
|
||||
self.size_database["__version__"] = dataloader_version
|
||||
|
||||
bad_count = 0
|
||||
for file in tqdm(file_list):
|
||||
@@ -456,25 +484,32 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
dataset_config=dataset_config,
|
||||
dataloader_transforms=self.transform,
|
||||
size_database=self.size_database,
|
||||
dataset_root=dataset_folder,
|
||||
)
|
||||
self.file_list.append(file_item)
|
||||
except Exception as e:
|
||||
print(traceback.format_exc())
|
||||
print(f"Error processing image: {file}")
|
||||
print(e)
|
||||
print_acc(traceback.format_exc())
|
||||
if self.is_video:
|
||||
print_acc(f"Error processing video: {file}")
|
||||
else:
|
||||
print_acc(f"Error processing image: {file}")
|
||||
print_acc(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}"
|
||||
|
||||
if self.is_video:
|
||||
print_acc(f" - Found {len(self.file_list)} videos")
|
||||
assert len(self.file_list) > 0, f"no videos found in {self.dataset_path}"
|
||||
else:
|
||||
print_acc(f" - Found {len(self.file_list)} images")
|
||||
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
|
||||
|
||||
# handle x axis flips
|
||||
if self.dataset_config.flip_x:
|
||||
print(" - adding x axis flips")
|
||||
print_acc(" - adding x axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the x axis
|
||||
@@ -484,7 +519,7 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
|
||||
# handle y axis flips
|
||||
if self.dataset_config.flip_y:
|
||||
print(" - adding y axis flips")
|
||||
print_acc(" - adding y axis flips")
|
||||
current_file_list = [x for x in self.file_list]
|
||||
for file_item in current_file_list:
|
||||
# create a copy that is flipped on the y axis
|
||||
@@ -493,8 +528,10 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
self.file_list.append(new_file_item)
|
||||
|
||||
if self.dataset_config.flip_x or self.dataset_config.flip_y:
|
||||
print(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
if self.is_video:
|
||||
print_acc(f" - Found {len(self.file_list)} videos after adding flips")
|
||||
else:
|
||||
print_acc(f" - Found {len(self.file_list)} images after adding flips")
|
||||
|
||||
self.setup_epoch()
|
||||
|
||||
@@ -522,7 +559,7 @@ class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, Capti
|
||||
return len(self.file_list)
|
||||
|
||||
def _get_single_item(self, index) -> 'FileItemDTO':
|
||||
file_item = copy.deepcopy(self.file_list[index])
|
||||
file_item: 'FileItemDTO' = copy.deepcopy(self.file_list[index])
|
||||
file_item.load_and_process_image(self.transform)
|
||||
file_item.load_caption(self.caption_dict)
|
||||
return file_item
|
||||
|
||||
@@ -2,6 +2,7 @@ import os
|
||||
import weakref
|
||||
from _weakref import ReferenceType
|
||||
from typing import TYPE_CHECKING, List, Union
|
||||
import cv2
|
||||
import torch
|
||||
import random
|
||||
|
||||
@@ -11,7 +12,7 @@ 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, ClipImageFileItemDTOMixin
|
||||
UnconditionalFileItemDTOMixin, ClipImageFileItemDTOMixin, InpaintControlFileItemDTOMixin
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -33,6 +34,7 @@ class FileItemDTO(
|
||||
CaptionProcessingDTOMixin,
|
||||
ImageProcessingDTOMixin,
|
||||
ControlFileItemDTOMixin,
|
||||
InpaintControlFileItemDTOMixin,
|
||||
ClipImageFileItemDTOMixin,
|
||||
MaskFileItemDTOMixin,
|
||||
AugmentationFileItemDTOMixin,
|
||||
@@ -43,20 +45,42 @@ class FileItemDTO(
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.path = kwargs.get('path', '')
|
||||
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
|
||||
self.is_video = self.dataset_config.num_frames > 1
|
||||
size_database = kwargs.get('size_database', {})
|
||||
filename = os.path.basename(self.path)
|
||||
if filename in size_database:
|
||||
w, h = size_database[filename]
|
||||
dataset_root = kwargs.get('dataset_root', None)
|
||||
if dataset_root is not None:
|
||||
# remove dataset root from path
|
||||
file_key = self.path.replace(dataset_root, '')
|
||||
else:
|
||||
file_key = os.path.basename(self.path)
|
||||
if file_key in size_database:
|
||||
w, h = size_database[file_key]
|
||||
elif self.is_video:
|
||||
# Open the video file
|
||||
video = cv2.VideoCapture(self.path)
|
||||
|
||||
# Check if video opened successfully
|
||||
if not video.isOpened():
|
||||
raise Exception(f"Error: Could not open video file {self.path}")
|
||||
|
||||
# Get width and height
|
||||
width = int(video.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(video.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
# Release the video capture object immediately
|
||||
video.release()
|
||||
size_database[file_key] = (width, height)
|
||||
else:
|
||||
# original method is significantly faster, but some images are read sideways. Not sure why. Do slow method for now.
|
||||
# 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)
|
||||
# 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))
|
||||
w, h = img.size
|
||||
size_database[file_key] = (w, h)
|
||||
self.width: int = w
|
||||
self.height: int = h
|
||||
self.dataloader_transforms = kwargs.get('dataloader_transforms', None)
|
||||
@@ -85,6 +109,7 @@ class FileItemDTO(
|
||||
self.tensor = None
|
||||
self.cleanup_latent()
|
||||
self.cleanup_control()
|
||||
self.cleanup_inpaint()
|
||||
self.cleanup_clip_image()
|
||||
self.cleanup_mask()
|
||||
self.cleanup_unconditional()
|
||||
@@ -131,6 +156,22 @@ class DataLoaderBatchDTO:
|
||||
else:
|
||||
control_tensors.append(x.control_tensor)
|
||||
self.control_tensor = torch.cat([x.unsqueeze(0) for x in control_tensors])
|
||||
|
||||
self.inpaint_tensor: Union[torch.Tensor, None] = None
|
||||
if any([x.inpaint_tensor is not None for x in self.file_items]):
|
||||
# find one to use as a base
|
||||
base_inpaint_tensor = None
|
||||
for x in self.file_items:
|
||||
if x.inpaint_tensor is not None:
|
||||
base_inpaint_tensor = x.inpaint_tensor
|
||||
break
|
||||
inpaint_tensors = []
|
||||
for x in self.file_items:
|
||||
if x.inpaint_tensor is None:
|
||||
inpaint_tensors.append(torch.zeros_like(base_inpaint_tensor))
|
||||
else:
|
||||
inpaint_tensors.append(x.inpaint_tensor)
|
||||
self.inpaint_tensor = torch.cat([x.unsqueeze(0) for x in inpaint_tensors])
|
||||
|
||||
self.loss_multiplier_list: List[float] = [x.loss_multiplier for x in self.file_items]
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import base64
|
||||
import glob
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
@@ -6,22 +7,26 @@ import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, List, Dict, Union
|
||||
import traceback
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from tqdm import tqdm
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection, SiglipImageProcessor
|
||||
|
||||
from toolkit.basic import flush, value_map
|
||||
from toolkit.buckets import get_bucket_for_image_size, get_resolution
|
||||
from toolkit.metadata import get_meta_for_safetensors
|
||||
from toolkit.models.pixtral_vision import PixtralVisionImagePreprocessorCompatible
|
||||
from toolkit.prompt_utils import inject_trigger_into_prompt
|
||||
from torchvision import transforms
|
||||
from PIL import Image, ImageFilter, ImageOps
|
||||
from PIL.ImageOps import exif_transpose
|
||||
import albumentations as A
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.accelerator import get_accelerator
|
||||
|
||||
from toolkit.train_tools import get_torch_dtype
|
||||
|
||||
@@ -30,6 +35,8 @@ if TYPE_CHECKING:
|
||||
from toolkit.data_transfer_object.data_loader import FileItemDTO
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
accelerator = get_accelerator()
|
||||
|
||||
# def get_associated_caption_from_img_path(img_path):
|
||||
# https://demo.albumentations.ai/
|
||||
class Augments:
|
||||
@@ -119,6 +126,9 @@ class CaptionMixin:
|
||||
prompt_path = path_no_ext + '.' + ext
|
||||
if os.path.exists(prompt_path):
|
||||
break
|
||||
|
||||
# allow folders to have a default prompt
|
||||
default_prompt_path = os.path.join(os.path.dirname(img_path), 'default.txt')
|
||||
|
||||
if os.path.exists(prompt_path):
|
||||
with open(prompt_path, 'r', encoding='utf-8') as f:
|
||||
@@ -129,6 +139,10 @@ class CaptionMixin:
|
||||
if 'caption' in prompt:
|
||||
prompt = prompt['caption']
|
||||
|
||||
prompt = clean_caption(prompt)
|
||||
elif os.path.exists(default_prompt_path):
|
||||
with open(default_prompt_path, 'r', encoding='utf-8') as f:
|
||||
prompt = f.read()
|
||||
prompt = clean_caption(prompt)
|
||||
else:
|
||||
prompt = ''
|
||||
@@ -254,7 +268,7 @@ class BucketsMixin:
|
||||
file_item.crop_y = int((file_item.scale_to_height - new_height) / 2)
|
||||
|
||||
if file_item.crop_y < 0 or file_item.crop_x < 0:
|
||||
print('debug')
|
||||
print_acc('debug')
|
||||
|
||||
# check if bucket exists, if not, create it
|
||||
bucket_key = f'{file_item.crop_width}x{file_item.crop_height}'
|
||||
@@ -266,10 +280,10 @@ class BucketsMixin:
|
||||
self.shuffle_buckets()
|
||||
self.build_batch_indices()
|
||||
if not quiet:
|
||||
print(f'Bucket sizes for {self.dataset_path}:')
|
||||
print_acc(f'Bucket sizes for {self.dataset_path}:')
|
||||
for key, bucket in self.buckets.items():
|
||||
print(f'{key}: {len(bucket.file_list_idx)} files')
|
||||
print(f'{len(self.buckets)} buckets made')
|
||||
print_acc(f'{key}: {len(bucket.file_list_idx)} files')
|
||||
print_acc(f'{len(self.buckets)} buckets made')
|
||||
|
||||
|
||||
class CaptionProcessingDTOMixin:
|
||||
@@ -417,16 +431,212 @@ class CaptionProcessingDTOMixin:
|
||||
|
||||
|
||||
class ImageProcessingDTOMixin:
|
||||
def load_and_process_video(
|
||||
self: 'FileItemDTO',
|
||||
transform: Union[None, transforms.Compose],
|
||||
only_load_latents=False
|
||||
):
|
||||
if self.is_latent_cached:
|
||||
raise Exception('Latent caching not supported for videos')
|
||||
|
||||
if self.augments is not None and len(self.augments) > 0:
|
||||
raise Exception('Augments not supported for videos')
|
||||
|
||||
if self.has_augmentations:
|
||||
raise Exception('Augmentations not supported for videos')
|
||||
|
||||
if not self.dataset_config.buckets:
|
||||
raise Exception('Buckets required for video processing')
|
||||
|
||||
try:
|
||||
# Use OpenCV to capture video frames
|
||||
cap = cv2.VideoCapture(self.path)
|
||||
|
||||
if not cap.isOpened():
|
||||
raise Exception(f"Failed to open video file: {self.path}")
|
||||
|
||||
# Get video properties
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
video_fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
|
||||
# Calculate the max valid frame index (accounting for zero-indexing)
|
||||
max_frame_index = total_frames - 1
|
||||
|
||||
# Only log video properties if in debug mode
|
||||
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
|
||||
print_acc(f"Video properties: {self.path}")
|
||||
print_acc(f" Total frames: {total_frames}")
|
||||
print_acc(f" Max valid frame index: {max_frame_index}")
|
||||
print_acc(f" FPS: {video_fps}")
|
||||
|
||||
frames_to_extract = []
|
||||
|
||||
# Always stretch/shrink to the requested number of frames if needed
|
||||
if self.dataset_config.shrink_video_to_frames or total_frames < self.dataset_config.num_frames:
|
||||
# Distribute frames evenly across the entire video
|
||||
interval = max_frame_index / (self.dataset_config.num_frames - 1) if self.dataset_config.num_frames > 1 else 0
|
||||
frames_to_extract = [min(int(round(i * interval)), max_frame_index) for i in range(self.dataset_config.num_frames)]
|
||||
else:
|
||||
# Calculate frame interval based on FPS ratio
|
||||
fps_ratio = video_fps / self.dataset_config.fps
|
||||
frame_interval = max(1, int(round(fps_ratio)))
|
||||
|
||||
# Calculate max consecutive frames we can extract at desired FPS
|
||||
max_consecutive_frames = (total_frames // frame_interval)
|
||||
|
||||
if max_consecutive_frames < self.dataset_config.num_frames:
|
||||
# Not enough frames at desired FPS, so stretch instead
|
||||
interval = max_frame_index / (self.dataset_config.num_frames - 1) if self.dataset_config.num_frames > 1 else 0
|
||||
frames_to_extract = [min(int(round(i * interval)), max_frame_index) for i in range(self.dataset_config.num_frames)]
|
||||
else:
|
||||
# Calculate max start frame to ensure we can get all num_frames
|
||||
max_start_frame = max_frame_index - ((self.dataset_config.num_frames - 1) * frame_interval)
|
||||
start_frame = random.randint(0, max(0, max_start_frame))
|
||||
|
||||
# Generate list of frames to extract
|
||||
frames_to_extract = [start_frame + (i * frame_interval) for i in range(self.dataset_config.num_frames)]
|
||||
|
||||
# Final safety check - ensure no frame exceeds max valid index
|
||||
frames_to_extract = [min(frame_idx, max_frame_index) for frame_idx in frames_to_extract]
|
||||
|
||||
# Only log frames to extract if in debug mode
|
||||
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
|
||||
print_acc(f" Frames to extract: {frames_to_extract}")
|
||||
|
||||
# Extract frames
|
||||
frames = []
|
||||
for frame_idx in frames_to_extract:
|
||||
# Safety check - ensure frame_idx is within bounds (silently fix)
|
||||
if frame_idx > max_frame_index:
|
||||
frame_idx = max_frame_index
|
||||
|
||||
# Set frame position
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)
|
||||
|
||||
# Silently verify position was set correctly (no warnings unless debug mode)
|
||||
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
|
||||
actual_pos = int(cap.get(cv2.CAP_PROP_POS_FRAMES))
|
||||
if actual_pos != frame_idx:
|
||||
print_acc(f"Warning: Failed to set exact frame position. Requested: {frame_idx}, Actual: {actual_pos}")
|
||||
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
# Try to provide more detailed error information
|
||||
actual_frame = int(cap.get(cv2.CAP_PROP_POS_FRAMES))
|
||||
frame_pos_info = f"Requested frame: {frame_idx}, Actual frame position: {actual_frame}"
|
||||
|
||||
# Try to read the next available frame as a fallback
|
||||
fallback_success = False
|
||||
for fallback_offset in [1, -1, 5, -5, 10, -10]:
|
||||
fallback_pos = max(0, min(frame_idx + fallback_offset, max_frame_index))
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, fallback_pos)
|
||||
fallback_ret, fallback_frame = cap.read()
|
||||
if fallback_ret:
|
||||
# Only log in debug mode
|
||||
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
|
||||
print_acc(f"Falling back to nearby frame {fallback_pos} instead of {frame_idx}")
|
||||
frame = fallback_frame
|
||||
fallback_success = True
|
||||
break
|
||||
else:
|
||||
# No fallback worked, raise a more detailed exception
|
||||
video_info = f"Video: {self.path}, Total frames: {total_frames}, FPS: {video_fps}"
|
||||
raise Exception(f"Failed to read frame {frame_idx} from video. {frame_pos_info}. {video_info}")
|
||||
|
||||
# Convert BGR to RGB
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
|
||||
# Convert to PIL Image
|
||||
img = Image.fromarray(frame)
|
||||
|
||||
# Apply the same processing as for single images
|
||||
img = img.convert('RGB')
|
||||
|
||||
if self.flip_x:
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
if self.flip_y:
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
# Apply bucketing
|
||||
img = img.resize((self.scale_to_width, self.scale_to_height), Image.BICUBIC)
|
||||
img = img.crop((
|
||||
self.crop_x,
|
||||
self.crop_y,
|
||||
self.crop_x + self.crop_width,
|
||||
self.crop_y + self.crop_height
|
||||
))
|
||||
|
||||
# Apply transform if provided
|
||||
if transform:
|
||||
img = transform(img)
|
||||
|
||||
frames.append(img)
|
||||
|
||||
# Release the video capture
|
||||
cap.release()
|
||||
|
||||
# Stack frames into tensor [frames, channels, height, width]
|
||||
self.tensor = torch.stack(frames)
|
||||
|
||||
# Only log success in debug mode
|
||||
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
|
||||
print_acc(f"Successfully loaded video with {len(frames)} frames: {self.path}")
|
||||
|
||||
except Exception as e:
|
||||
# Print full traceback
|
||||
traceback.print_exc()
|
||||
|
||||
# Provide more context about the error
|
||||
error_msg = str(e)
|
||||
try:
|
||||
if 'Failed to read frame' in error_msg and cap is not None:
|
||||
# Try to get more info about the video that failed
|
||||
cap_status = "Opened" if cap.isOpened() else "Closed"
|
||||
current_pos = int(cap.get(cv2.CAP_PROP_POS_FRAMES)) if cap.isOpened() else "Unknown"
|
||||
reported_total = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) if cap.isOpened() else "Unknown"
|
||||
|
||||
print_acc(f"Video details when error occurred:")
|
||||
print_acc(f" Cap status: {cap_status}")
|
||||
print_acc(f" Current position: {current_pos}")
|
||||
print_acc(f" Reported total frames: {reported_total}")
|
||||
|
||||
# Try to verify if the video is corrupted
|
||||
if cap.isOpened():
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, 0) # Go to start
|
||||
start_ret, _ = cap.read()
|
||||
|
||||
# Try to read the last frame to check if it's accessible
|
||||
if reported_total > 0:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, reported_total - 1)
|
||||
end_ret, _ = cap.read()
|
||||
print_acc(f" Can read first frame: {start_ret}, Can read last frame: {end_ret}")
|
||||
|
||||
# Close the cap if it's still open
|
||||
cap.release()
|
||||
except Exception as debug_err:
|
||||
print_acc(f"Error during error diagnosis: {debug_err}")
|
||||
|
||||
print_acc(f"Error: {error_msg}")
|
||||
print_acc(f"Error loading video: {self.path}")
|
||||
|
||||
# Re-raise with more detailed information
|
||||
raise Exception(f"Video loading error ({self.path}): {error_msg}") from e
|
||||
|
||||
def load_and_process_image(
|
||||
self: 'FileItemDTO',
|
||||
transform: Union[None, transforms.Compose],
|
||||
only_load_latents=False
|
||||
):
|
||||
if self.dataset_config.num_frames > 1:
|
||||
self.load_and_process_video(transform, only_load_latents)
|
||||
return
|
||||
# if we are caching latents, just do that
|
||||
if self.is_latent_cached:
|
||||
self.get_latent()
|
||||
if self.has_control_image:
|
||||
self.load_control_image()
|
||||
if self.has_inpaint_image:
|
||||
self.load_inpaint_image()
|
||||
if self.has_clip_image:
|
||||
self.load_clip_image()
|
||||
if self.has_mask_image:
|
||||
@@ -438,8 +648,8 @@ class ImageProcessingDTOMixin:
|
||||
img = Image.open(self.path)
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.path}")
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {self.path}")
|
||||
|
||||
if self.use_alpha_as_mask:
|
||||
# we do this to make sure it does not replace the alpha with another color
|
||||
@@ -453,11 +663,11 @@ class ImageProcessingDTOMixin:
|
||||
w, h = img.size
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
print(
|
||||
print_acc(
|
||||
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
|
||||
print(
|
||||
print_acc(
|
||||
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:
|
||||
@@ -473,7 +683,7 @@ class ImageProcessingDTOMixin:
|
||||
# crop to x_crop, y_crop, x_crop + crop_width, y_crop + crop_height
|
||||
if img.width < self.crop_x + self.crop_width or img.height < self.crop_y + self.crop_height:
|
||||
# todo look into this. This still happens sometimes
|
||||
print('size mismatch')
|
||||
print_acc('size mismatch')
|
||||
img = img.crop((
|
||||
self.crop_x,
|
||||
self.crop_y,
|
||||
@@ -492,7 +702,7 @@ class ImageProcessingDTOMixin:
|
||||
if self.dataset_config.random_crop:
|
||||
if self.dataset_config.random_scale and min_img_size > self.dataset_config.resolution:
|
||||
if min_img_size < self.dataset_config.resolution:
|
||||
print(
|
||||
print_acc(
|
||||
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.dataset_config.resolution}, image file={self.path}")
|
||||
scale_size = self.dataset_config.resolution
|
||||
else:
|
||||
@@ -522,6 +732,8 @@ class ImageProcessingDTOMixin:
|
||||
if not only_load_latents:
|
||||
if self.has_control_image:
|
||||
self.load_control_image()
|
||||
if self.has_inpaint_image:
|
||||
self.load_inpaint_image()
|
||||
if self.has_clip_image:
|
||||
self.load_clip_image()
|
||||
if self.has_mask_image:
|
||||
@@ -530,43 +742,38 @@ class ImageProcessingDTOMixin:
|
||||
self.load_unconditional_image()
|
||||
|
||||
|
||||
class ControlFileItemDTOMixin:
|
||||
class InpaintControlFileItemDTOMixin:
|
||||
def __init__(self: 'FileItemDTO', *args, **kwargs):
|
||||
if hasattr(super(), '__init__'):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.has_control_image = False
|
||||
self.control_path: Union[str, None] = None
|
||||
self.control_tensor: Union[torch.Tensor, None] = None
|
||||
self.has_inpaint_image = False
|
||||
self.inpaint_path: Union[str, None] = None
|
||||
self.inpaint_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:
|
||||
if dataset_config.inpaint_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
|
||||
inpaint_path = dataset_config.inpaint_path
|
||||
# we are using control images
|
||||
img_path = kwargs.get('path', None)
|
||||
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
img_ext_list = ['.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(control_path, file_name_no_ext + ext)):
|
||||
self.control_path = os.path.join(control_path, file_name_no_ext + ext)
|
||||
self.has_control_image = True
|
||||
p = os.path.join(inpaint_path, file_name_no_ext + ext)
|
||||
if os.path.exists(p):
|
||||
self.inpaint_path = p
|
||||
self.has_inpaint_image = True
|
||||
break
|
||||
|
||||
def load_control_image(self: 'FileItemDTO'):
|
||||
|
||||
def load_inpaint_image(self: 'FileItemDTO'):
|
||||
try:
|
||||
img = Image.open(self.control_path).convert('RGB')
|
||||
# image must have alpha channel for inpaint
|
||||
img = Image.open(self.inpaint_path)
|
||||
# make sure has aplha
|
||||
if img.mode != 'RGBA':
|
||||
raise ValueError(f"Image must have alpha channel for inpaint: {self.inpaint_path}")
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.control_path}")
|
||||
|
||||
if self.full_size_control_images:
|
||||
# we just scale them to 512x512:
|
||||
w, h = img.size
|
||||
img = img.resize((512, 512), Image.BICUBIC)
|
||||
|
||||
else:
|
||||
|
||||
w, h = img.size
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
@@ -596,14 +803,127 @@ class ControlFileItemDTOMixin:
|
||||
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)
|
||||
raise Exception("Inpaint images not supported for non-bucket datasets")
|
||||
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
if self.aug_replay_spatial_transforms:
|
||||
tensor = self.augment_spatial_control(img, transform=transform)
|
||||
else:
|
||||
tensor = transform(img)
|
||||
|
||||
# is 0 to 1 with alpha
|
||||
self.inpaint_tensor = tensor
|
||||
|
||||
except Exception as e:
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {self.inpaint_path}")
|
||||
|
||||
|
||||
def cleanup_inpaint(self: 'FileItemDTO'):
|
||||
self.inpaint_tensor = None
|
||||
|
||||
|
||||
class ControlFileItemDTOMixin:
|
||||
def __init__(self: 'FileItemDTO', *args, **kwargs):
|
||||
if hasattr(super(), '__init__'):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.has_control_image = False
|
||||
self.control_path: Union[str, List[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_list = dataset_config.control_path
|
||||
if not isinstance(control_path_list, list):
|
||||
control_path_list = [control_path_list]
|
||||
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']
|
||||
file_name_no_ext = os.path.splitext(os.path.basename(img_path))[0]
|
||||
|
||||
found_control_images = []
|
||||
for control_path in control_path_list:
|
||||
for ext in img_ext_list:
|
||||
if os.path.exists(os.path.join(control_path, file_name_no_ext + ext)):
|
||||
found_control_images.append(os.path.join(control_path, file_name_no_ext + ext))
|
||||
self.has_control_image = True
|
||||
break
|
||||
self.control_path = found_control_images
|
||||
if len(self.control_path) == 0:
|
||||
self.control_path = None
|
||||
elif len(self.control_path) == 1:
|
||||
# only do one
|
||||
self.control_path = self.control_path[0]
|
||||
|
||||
def load_control_image(self: 'FileItemDTO'):
|
||||
control_tensors = []
|
||||
control_path_list = self.control_path
|
||||
if not isinstance(self.control_path, list):
|
||||
control_path_list = [self.control_path]
|
||||
|
||||
for control_path in control_path_list:
|
||||
try:
|
||||
img = Image.open(control_path).convert('RGB')
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {control_path}")
|
||||
|
||||
if not self.full_size_control_images:
|
||||
# we just scale them to 512x512:
|
||||
w, h = img.size
|
||||
img = img.resize((512, 512), Image.BICUBIC)
|
||||
|
||||
else:
|
||||
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 = 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:
|
||||
tensor = self.augment_spatial_control(img, transform=transform)
|
||||
else:
|
||||
tensor = transform(img)
|
||||
control_tensors.append(tensor)
|
||||
|
||||
if len(control_tensors) == 0:
|
||||
self.control_tensor = None
|
||||
elif len(control_tensors) == 1:
|
||||
self.control_tensor = control_tensors[0]
|
||||
else:
|
||||
self.control_tensor = transform(img)
|
||||
self.control_tensor = torch.stack(control_tensors, dim=0)
|
||||
|
||||
def cleanup_control(self: 'FileItemDTO'):
|
||||
self.control_tensor = None
|
||||
@@ -629,11 +949,12 @@ class ClipImageFileItemDTOMixin:
|
||||
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:
|
||||
if dataset_config.clip_image_path is not None or dataset_config.clip_image_from_same_folder:
|
||||
# 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
|
||||
if dataset_config.clip_image_path is not None:
|
||||
# find the control image path
|
||||
clip_image_path = dataset_config.clip_image_path
|
||||
# we are using control images
|
||||
@@ -645,7 +966,11 @@ class ClipImageFileItemDTOMixin:
|
||||
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()
|
||||
|
||||
if dataset_config.clip_image_from_same_folder:
|
||||
# assume we have one. We will pull it on load.
|
||||
self.has_clip_image = True
|
||||
self.build_clip_imag_augmentation_transform()
|
||||
|
||||
def build_clip_imag_augmentation_transform(self: 'FileItemDTO'):
|
||||
@@ -731,8 +1056,29 @@ class ClipImageFileItemDTOMixin:
|
||||
self._clip_vision_embeddings_path = os.path.join(latent_dir, f'{filename_no_ext}_{hash_str}.safetensors')
|
||||
|
||||
return self._clip_vision_embeddings_path
|
||||
|
||||
def get_new_clip_image_path(self: 'FileItemDTO'):
|
||||
if self.dataset_config.clip_image_from_same_folder:
|
||||
# randomly grab an image path from the same folder
|
||||
pool_folder = os.path.dirname(self.path)
|
||||
# find all images in the folder
|
||||
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
|
||||
img_files = []
|
||||
for ext in img_ext_list:
|
||||
img_files += glob.glob(os.path.join(pool_folder, f'*{ext}'))
|
||||
# remove the current image if len is greater than 1
|
||||
if len(img_files) > 1:
|
||||
img_files.remove(self.path)
|
||||
# randomly grab one
|
||||
return random.choice(img_files)
|
||||
else:
|
||||
return self.clip_image_path
|
||||
|
||||
def load_clip_image(self: 'FileItemDTO'):
|
||||
is_dynamic_size_and_aspect = isinstance(self.clip_image_processor, PixtralVisionImagePreprocessorCompatible) or \
|
||||
isinstance(self.clip_image_processor, SiglipImageProcessor)
|
||||
if self.clip_image_processor is None:
|
||||
is_dynamic_size_and_aspect = True # serving it raw
|
||||
if self.is_vision_clip_cached:
|
||||
self.clip_image_embeds = load_file(self.get_clip_vision_embeddings_path())
|
||||
|
||||
@@ -742,14 +1088,15 @@ class ClipImageFileItemDTOMixin:
|
||||
self.clip_image_embeds_unconditional = load_file(unconditional_path)
|
||||
|
||||
return
|
||||
clip_image_path = self.get_new_clip_image_path()
|
||||
try:
|
||||
img = Image.open(self.clip_image_path).convert('RGB')
|
||||
img = Image.open(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}")
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {clip_image_path}")
|
||||
|
||||
img = img.convert('RGB')
|
||||
|
||||
@@ -759,8 +1106,10 @@ class ClipImageFileItemDTOMixin:
|
||||
if self.flip_y:
|
||||
# do a flip
|
||||
img = img.transpose(Image.FLIP_TOP_BOTTOM)
|
||||
|
||||
if img.width != img.height:
|
||||
|
||||
if is_dynamic_size_and_aspect:
|
||||
pass # let the image processor handle it
|
||||
elif img.width != img.height:
|
||||
min_size = min(img.width, img.height)
|
||||
if self.dataset_config.square_crop:
|
||||
# center crop to a square
|
||||
@@ -945,8 +1294,8 @@ class MaskFileItemDTOMixin:
|
||||
img = Image.open(self.mask_path)
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.mask_path}")
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {self.mask_path}")
|
||||
|
||||
if self.use_alpha_as_mask:
|
||||
# pipeline expectws an rgb image so we need to put alpha in all channels
|
||||
@@ -963,11 +1312,11 @@ class MaskFileItemDTOMixin:
|
||||
fix_size = False
|
||||
if w > h and self.scale_to_width < self.scale_to_height:
|
||||
# throw error, they should match
|
||||
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}")
|
||||
print_acc(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
|
||||
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}")
|
||||
print_acc(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:
|
||||
@@ -1049,8 +1398,8 @@ class UnconditionalFileItemDTOMixin:
|
||||
img = Image.open(self.unconditional_path)
|
||||
img = exif_transpose(img)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
print(f"Error loading image: {self.mask_path}")
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error loading image: {self.mask_path}")
|
||||
|
||||
img = img.convert('RGB')
|
||||
w, h = img.size
|
||||
@@ -1130,9 +1479,9 @@ class PoiFileItemDTOMixin:
|
||||
with open(caption_path, 'r', encoding='utf-8') as f:
|
||||
json_data = json.load(f)
|
||||
if 'poi' not in json_data:
|
||||
print(f"Warning: poi not found in caption file: {caption_path}")
|
||||
print_acc(f"Warning: poi not found in caption file: {caption_path}")
|
||||
if self.poi not in json_data['poi']:
|
||||
print(f"Warning: poi not found in caption file: {caption_path}")
|
||||
print_acc(f"Warning: poi not found in caption file: {caption_path}")
|
||||
# poi has, x, y, width, height
|
||||
# do full image if no poi
|
||||
self.poi_x = 0
|
||||
@@ -1206,8 +1555,8 @@ class PoiFileItemDTOMixin:
|
||||
# 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}")
|
||||
print_acc(f"Error: {e}")
|
||||
print_acc(f"Error getting resolution: {self.path}")
|
||||
raise e
|
||||
return False
|
||||
if current_resolution >= self.dataset_config.resolution:
|
||||
@@ -1216,7 +1565,7 @@ class PoiFileItemDTOMixin:
|
||||
else:
|
||||
num_loops += 1
|
||||
if num_loops > 100:
|
||||
print(
|
||||
print_acc(
|
||||
f"Warning: poi bucketing looped too many times. This should not happen. Please report this issue.")
|
||||
return False
|
||||
|
||||
@@ -1243,7 +1592,7 @@ class PoiFileItemDTOMixin:
|
||||
|
||||
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')
|
||||
print_acc('size mismatch')
|
||||
|
||||
return True
|
||||
|
||||
@@ -1337,88 +1686,91 @@ class LatentCachingMixin:
|
||||
self.latent_cache = {}
|
||||
|
||||
def cache_latents_all_latents(self: 'AiToolkitDataset'):
|
||||
print(f"Caching latents for {self.dataset_path}")
|
||||
# cache all latents to disk
|
||||
to_disk = self.is_caching_latents_to_disk
|
||||
to_memory = self.is_caching_latents_to_memory
|
||||
if self.dataset_config.num_frames > 1:
|
||||
raise Exception("Error: caching latents is not supported for multi-frame datasets")
|
||||
with accelerator.main_process_first():
|
||||
print_acc(f"Caching latents for {self.dataset_path}")
|
||||
# cache all latents to disk
|
||||
to_disk = self.is_caching_latents_to_disk
|
||||
to_memory = self.is_caching_latents_to_memory
|
||||
|
||||
if to_disk:
|
||||
print(" - Saving latents to disk")
|
||||
if to_memory:
|
||||
print(" - Keeping latents in memory")
|
||||
# move sd items to cpu except for vae
|
||||
self.sd.set_device_state_preset('cache_latents')
|
||||
if to_disk:
|
||||
print_acc(" - Saving latents to disk")
|
||||
if to_memory:
|
||||
print_acc(" - Keeping latents in memory")
|
||||
# move sd items to cpu except for vae
|
||||
self.sd.set_device_state_preset('cache_latents')
|
||||
|
||||
# use tqdm to show progress
|
||||
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.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'
|
||||
file_item.is_caching_to_disk = to_disk
|
||||
file_item.is_caching_to_memory = to_memory
|
||||
file_item.latent_load_device = self.sd.device
|
||||
# use tqdm to show progress
|
||||
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.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 = self.sd.model_config.arch
|
||||
file_item.is_caching_to_disk = to_disk
|
||||
file_item.is_caching_to_memory = to_memory
|
||||
file_item.latent_load_device = self.sd.device
|
||||
|
||||
latent_path = file_item.get_latent_path(recalculate=True)
|
||||
# check if it is saved to disk already
|
||||
if os.path.exists(latent_path):
|
||||
if to_memory:
|
||||
# load it into memory
|
||||
state_dict = load_file(latent_path, device='cpu')
|
||||
file_item._encoded_latent = state_dict['latent'].to('cpu', dtype=self.sd.torch_dtype)
|
||||
else:
|
||||
# not saved to disk, calculate
|
||||
# load the image first
|
||||
file_item.load_and_process_image(self.transform, only_load_latents=True)
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
# add batch dimension
|
||||
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([
|
||||
('latent', latent.clone().detach().cpu()),
|
||||
])
|
||||
# metadata
|
||||
meta = get_meta_for_safetensors(file_item.get_latent_info_dict())
|
||||
os.makedirs(os.path.dirname(latent_path), exist_ok=True)
|
||||
save_file(state_dict, latent_path, metadata=meta)
|
||||
latent_path = file_item.get_latent_path(recalculate=True)
|
||||
# check if it is saved to disk already
|
||||
if os.path.exists(latent_path):
|
||||
if to_memory:
|
||||
# load it into memory
|
||||
state_dict = load_file(latent_path, device='cpu')
|
||||
file_item._encoded_latent = state_dict['latent'].to('cpu', dtype=self.sd.torch_dtype)
|
||||
else:
|
||||
# not saved to disk, calculate
|
||||
# load the image first
|
||||
file_item.load_and_process_image(self.transform, only_load_latents=True)
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
# add batch dimension
|
||||
try:
|
||||
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
latent = self.sd.encode_images(imgs).squeeze(0)
|
||||
except Exception as e:
|
||||
print_acc(f"Error processing image: {file_item.path}")
|
||||
print_acc(f"Error: {str(e)}")
|
||||
raise e
|
||||
# save_latent
|
||||
if to_disk:
|
||||
state_dict = OrderedDict([
|
||||
('latent', latent.clone().detach().cpu()),
|
||||
])
|
||||
# metadata
|
||||
meta = get_meta_for_safetensors(file_item.get_latent_info_dict())
|
||||
os.makedirs(os.path.dirname(latent_path), exist_ok=True)
|
||||
save_file(state_dict, latent_path, metadata=meta)
|
||||
|
||||
if to_memory:
|
||||
# keep it in memory
|
||||
file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if to_memory:
|
||||
# keep it in memory
|
||||
file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
|
||||
del imgs
|
||||
del latent
|
||||
del file_item.tensor
|
||||
del imgs
|
||||
del latent
|
||||
del file_item.tensor
|
||||
|
||||
# flush(garbage_collect=False)
|
||||
file_item.is_latent_cached = True
|
||||
i += 1
|
||||
# flush every 100
|
||||
# if i % 100 == 0:
|
||||
# flush()
|
||||
# flush(garbage_collect=False)
|
||||
file_item.is_latent_cached = True
|
||||
i += 1
|
||||
# flush every 100
|
||||
# if i % 100 == 0:
|
||||
# flush()
|
||||
|
||||
# restore device state
|
||||
self.sd.restore_device_state()
|
||||
# restore device state
|
||||
self.sd.restore_device_state()
|
||||
|
||||
|
||||
class CLIPCachingMixin:
|
||||
@@ -1433,9 +1785,9 @@ class CLIPCachingMixin:
|
||||
if not self.is_caching_clip_vision_to_disk:
|
||||
return
|
||||
with torch.no_grad():
|
||||
print(f"Caching clip vision for {self.dataset_path}")
|
||||
print_acc(f"Caching clip vision for {self.dataset_path}")
|
||||
|
||||
print(" - Saving clip to disk")
|
||||
print_acc(" - Saving clip to disk")
|
||||
# move sd items to cpu except for vae
|
||||
self.sd.set_device_state_preset('cache_clip')
|
||||
|
||||
@@ -1476,7 +1828,7 @@ class CLIPCachingMixin:
|
||||
self.clip_vision_num_unconditional_cache = 1
|
||||
|
||||
# cache unconditionals
|
||||
print(f" - Caching {self.clip_vision_num_unconditional_cache} unconditional clip vision to disk")
|
||||
print_acc(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 = []
|
||||
|
||||
88
toolkit/dequantize.py
Normal file
88
toolkit/dequantize.py
Normal file
@@ -0,0 +1,88 @@
|
||||
|
||||
|
||||
from functools import partial
|
||||
from optimum.quanto.tensor import QTensor
|
||||
import torch
|
||||
|
||||
|
||||
def hacked_state_dict(self, *args, **kwargs):
|
||||
orig_state_dict = self.orig_state_dict(*args, **kwargs)
|
||||
new_state_dict = {}
|
||||
for key, value in orig_state_dict.items():
|
||||
if key.endswith("._scale"):
|
||||
continue
|
||||
if key.endswith(".input_scale"):
|
||||
continue
|
||||
if key.endswith(".output_scale"):
|
||||
continue
|
||||
if key.endswith("._data"):
|
||||
key = key[:-6]
|
||||
scale = orig_state_dict[key + "._scale"]
|
||||
# scale is the original dtype
|
||||
dtype = scale.dtype
|
||||
scale = scale.float()
|
||||
value = value.float()
|
||||
dequantized = value * scale
|
||||
|
||||
# handle input and output scaling if they exist
|
||||
input_scale = orig_state_dict.get(key + ".input_scale")
|
||||
|
||||
if input_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if input_scale.item() != 1.0:
|
||||
raise ValueError("Input scale is not 1.0, cannot dequantize")
|
||||
|
||||
output_scale = orig_state_dict.get(key + ".output_scale")
|
||||
|
||||
if output_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if output_scale.item() != 1.0:
|
||||
raise ValueError("Output scale is not 1.0, cannot dequantize")
|
||||
|
||||
new_state_dict[key] = dequantized.to('cpu', dtype=dtype)
|
||||
else:
|
||||
new_state_dict[key] = value
|
||||
return new_state_dict
|
||||
|
||||
# hacks the state dict so we can dequantize before saving
|
||||
def patch_dequantization_on_save(model):
|
||||
model.orig_state_dict = model.state_dict
|
||||
model.state_dict = partial(hacked_state_dict, model)
|
||||
|
||||
|
||||
def dequantize_parameter(module: torch.nn.Module, param_name: str) -> bool:
|
||||
"""
|
||||
Convert a quantized parameter back to a regular Parameter with floating point values.
|
||||
|
||||
Args:
|
||||
module: The module containing the parameter to unquantize
|
||||
param_name: Name of the parameter to unquantize (e.g., 'weight', 'bias')
|
||||
|
||||
Returns:
|
||||
bool: True if parameter was unquantized, False if it was already unquantized
|
||||
"""
|
||||
|
||||
# Check if the parameter exists
|
||||
if not hasattr(module, param_name):
|
||||
raise AttributeError(f"Module has no parameter named '{param_name}'")
|
||||
|
||||
param = getattr(module, param_name)
|
||||
|
||||
# If it's not a parameter or not quantized, nothing to do
|
||||
if not isinstance(param, torch.nn.Parameter):
|
||||
raise TypeError(f"'{param_name}' is not a Parameter")
|
||||
if not isinstance(param, QTensor):
|
||||
return False
|
||||
|
||||
# Convert to float tensor while preserving device and requires_grad
|
||||
with torch.no_grad():
|
||||
float_tensor = param.float()
|
||||
new_param = torch.nn.Parameter(
|
||||
float_tensor,
|
||||
requires_grad=param.requires_grad
|
||||
)
|
||||
|
||||
# Replace the parameter
|
||||
setattr(module, param_name, new_param)
|
||||
|
||||
return True
|
||||
@@ -5,6 +5,7 @@ from typing import Iterable, Optional
|
||||
import weakref
|
||||
import copy
|
||||
import contextlib
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic
|
||||
|
||||
import torch
|
||||
|
||||
@@ -43,9 +44,10 @@ class ExponentialMovingAverage:
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter] = None,
|
||||
decay: float = 0.995,
|
||||
use_num_updates: bool = True,
|
||||
use_num_updates: bool = False,
|
||||
# feeds back the decat to the parameter
|
||||
use_feedback: bool = False
|
||||
use_feedback: bool = False,
|
||||
param_multiplier: float = 1.0
|
||||
):
|
||||
if parameters is None:
|
||||
raise ValueError("parameters must be provided")
|
||||
@@ -54,6 +56,7 @@ class ExponentialMovingAverage:
|
||||
self.decay = decay
|
||||
self.num_updates = 0 if use_num_updates else None
|
||||
self.use_feedback = use_feedback
|
||||
self.param_multiplier = param_multiplier
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [
|
||||
p.clone().detach()
|
||||
@@ -121,13 +124,32 @@ class ExponentialMovingAverage:
|
||||
one_minus_decay = 1.0 - decay
|
||||
with torch.no_grad():
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
tmp = (s_param - param)
|
||||
s_param_float = s_param.float()
|
||||
if s_param.dtype != torch.float32:
|
||||
s_param_float = s_param_float.to(torch.float32)
|
||||
param_float = param
|
||||
if param.dtype != torch.float32:
|
||||
param_float = param_float.to(torch.float32)
|
||||
tmp = (s_param_float - param_float)
|
||||
# tmp will be a new tensor so we can do in-place
|
||||
tmp.mul_(one_minus_decay)
|
||||
s_param.sub_(tmp)
|
||||
|
||||
s_param_float.sub_(tmp)
|
||||
|
||||
update_param = False
|
||||
if self.use_feedback:
|
||||
param.add_(tmp)
|
||||
param_float.add_(tmp)
|
||||
update_param = True
|
||||
|
||||
if self.param_multiplier != 1.0:
|
||||
param_float.mul_(self.param_multiplier)
|
||||
update_param = True
|
||||
|
||||
if s_param.dtype != torch.float32:
|
||||
copy_stochastic(s_param, s_param_float)
|
||||
|
||||
if update_param and param.dtype != torch.float32:
|
||||
copy_stochastic(param, param_float)
|
||||
|
||||
|
||||
def copy_to(
|
||||
self,
|
||||
|
||||
@@ -6,6 +6,7 @@ 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
|
||||
from toolkit.config_modules import TrainConfig
|
||||
|
||||
GuidanceType = Literal["targeted", "polarity", "targeted_polarity", "direct"]
|
||||
|
||||
@@ -23,8 +24,12 @@ def get_differential_mask(
|
||||
):
|
||||
# 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]
|
||||
if len(differential_mask.shape) == 4:
|
||||
max_differential = \
|
||||
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
|
||||
elif len(differential_mask.shape) == 5:
|
||||
max_differential = \
|
||||
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0].max(dim=4, keepdim=True)[0]
|
||||
differential_scaler = 1.0 / max_differential
|
||||
differential_mask = differential_mask * differential_scaler
|
||||
|
||||
@@ -407,6 +412,7 @@ def get_guided_loss_polarity(
|
||||
batch: 'DataLoaderBatchDTO',
|
||||
noise: torch.Tensor,
|
||||
sd: 'StableDiffusion',
|
||||
train_config: 'TrainConfig',
|
||||
scaler=None,
|
||||
**kwargs
|
||||
):
|
||||
@@ -423,8 +429,22 @@ def get_guided_loss_polarity(
|
||||
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)
|
||||
linear_timesteps = any([
|
||||
train_config.linear_timesteps,
|
||||
train_config.linear_timesteps2,
|
||||
train_config.timestep_type == 'linear',
|
||||
])
|
||||
|
||||
timestep_type = 'linear' if linear_timesteps else None
|
||||
if timestep_type is None:
|
||||
timestep_type = train_config.timestep_type
|
||||
|
||||
sd.noise_scheduler.set_train_timesteps(
|
||||
1000,
|
||||
device=device,
|
||||
timestep_type=timestep_type,
|
||||
latents=conditional_latents
|
||||
)
|
||||
target_pos = (noise - conditional_latents).detach()
|
||||
target_neg = (noise - unconditional_latents).detach()
|
||||
|
||||
@@ -481,11 +501,6 @@ def get_guided_loss_polarity(
|
||||
|
||||
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:
|
||||
@@ -592,6 +607,105 @@ def get_guided_tnt(
|
||||
|
||||
return loss
|
||||
|
||||
def targeted_flow_guidance(
|
||||
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,
|
||||
train_config=None,
|
||||
**kwargs
|
||||
):
|
||||
if not sd.is_flow_matching:
|
||||
raise ValueError("targeted_flow only works on flow matching models")
|
||||
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()
|
||||
|
||||
# get a mask on the differential of the latents
|
||||
# this will be scaled from 0.0-1.0 with 1.0 being the largest differential
|
||||
abs_differential_mask = get_differential_mask(
|
||||
conditional_latents,
|
||||
unconditional_latents,
|
||||
gradient=True
|
||||
)
|
||||
|
||||
# get noisy latents for both conditional and unconditional predictions
|
||||
unconditional_noisy_latents = sd.add_noise(
|
||||
unconditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
conditional_noisy_latents = sd.add_noise(
|
||||
conditional_latents,
|
||||
noise,
|
||||
timesteps
|
||||
).detach()
|
||||
|
||||
# disable the lora to get a baseline prediction
|
||||
sd.network.is_active = False
|
||||
sd.unet.eval()
|
||||
|
||||
# get a baseline prediction of the model knowledge without the lora network
|
||||
# we do this with the unconditional noisy latents
|
||||
baseline_prediction = 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
|
||||
).detach()
|
||||
|
||||
# This is our normal flowmatching target
|
||||
# target = noise - latents
|
||||
# we need to target the baseline noise but with our conditional latents
|
||||
# to do this we first have to determine the baseline_prediction noise by reversing the flowmatching target
|
||||
baseline_predicted_noise = baseline_prediction + unconditional_latents
|
||||
|
||||
# baseline_predicted_noise is now the noise prediction our model would make with a the unconditional image.
|
||||
# we use this as our new noise target to preserve the existing knowledge of the image.
|
||||
# we apply a mask to this noise to only allow the differential of the conditional latents to be learned
|
||||
baseline_predicted_noise = (1 - abs_differential_mask) * baseline_predicted_noise
|
||||
masked_noise = abs_differential_mask * noise
|
||||
target_noise = masked_noise + baseline_predicted_noise
|
||||
|
||||
# compute our new target prediction using our current knowledge noise with our conditional latents
|
||||
# this makes it so the only new information is the differential of our conditional and unconditional latents
|
||||
# forcing the network to preserve existing knowledge, but learn only our changes
|
||||
target_pred = (target_noise - conditional_latents).detach()
|
||||
|
||||
# make a prediction with the lora network active
|
||||
sd.unet.train()
|
||||
sd.network.is_active = True
|
||||
sd.network.multiplier = network_weight_list
|
||||
prediction = sd.predict_noise(
|
||||
latents=conditional_noisy_latents.to(device, dtype=dtype).detach(),
|
||||
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
|
||||
timestep=timesteps,
|
||||
guidance_scale=1.0,
|
||||
**pred_kwargs
|
||||
)
|
||||
|
||||
# target our baseline + diffirential noise target
|
||||
pred_loss = torch.nn.functional.mse_loss(
|
||||
prediction.float(),
|
||||
target_pred.float()
|
||||
)
|
||||
|
||||
return pred_loss
|
||||
|
||||
|
||||
# this processes all guidance losses based on the batch information
|
||||
@@ -609,6 +723,7 @@ def get_guidance_loss(
|
||||
mask_multiplier=None,
|
||||
prior_pred=None,
|
||||
scaler=None,
|
||||
train_config=None,
|
||||
**kwargs
|
||||
):
|
||||
# TODO add others and process individual batch items separately
|
||||
@@ -641,6 +756,7 @@ def get_guidance_loss(
|
||||
noise,
|
||||
sd,
|
||||
scaler=scaler,
|
||||
train_config=train_config,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "tnt":
|
||||
@@ -689,5 +805,23 @@ def get_guidance_loss(
|
||||
prior_pred=prior_pred,
|
||||
**kwargs
|
||||
)
|
||||
elif guidance_type == "targeted_flow":
|
||||
return targeted_flow_guidance(
|
||||
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,
|
||||
scaler=scaler,
|
||||
train_config=train_config,
|
||||
**kwargs
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Guidance type {guidance_type} is not implemented")
|
||||
|
||||
@@ -12,6 +12,7 @@ import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import AutoencoderTiny
|
||||
from PIL import Image as PILImage
|
||||
|
||||
FILE_UNKNOWN = "Sorry, don't know how to get size for this file."
|
||||
|
||||
@@ -480,7 +481,26 @@ def show_tensors(imgs: torch.Tensor, name='AI Toolkit'):
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
|
||||
show_img(img_numpy[0], name=name)
|
||||
|
||||
def save_tensors(imgs: torch.Tensor, path='output.png'):
|
||||
if len(imgs.shape) == 5 and imgs.shape[0] == 1:
|
||||
imgs = imgs.squeeze(0)
|
||||
if len(imgs.shape) == 4:
|
||||
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
|
||||
else:
|
||||
img_list = [imgs]
|
||||
|
||||
img = torch.cat(img_list, dim=3)
|
||||
img = img / 2 + 0.5
|
||||
img_numpy = img.to(torch.float32).detach().cpu().numpy()
|
||||
img_numpy = np.clip(img_numpy, 0, 1) * 255
|
||||
img_numpy = img_numpy.transpose(0, 2, 3, 1)
|
||||
img_numpy = img_numpy.astype(np.uint8)
|
||||
# concat images to one
|
||||
img_numpy = np.concatenate(img_numpy, axis=1)
|
||||
# conver to pil
|
||||
img_pil = PILImage.fromarray(img_numpy)
|
||||
img_pil.save(path)
|
||||
|
||||
def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit'):
|
||||
if vae.device == 'cpu':
|
||||
|
||||
@@ -269,16 +269,7 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
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]
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
@@ -297,7 +288,44 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# will be none if disabled
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `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:
|
||||
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)
|
||||
|
||||
# begin ip adapter
|
||||
if not is_active:
|
||||
ip_hidden_states = None
|
||||
else:
|
||||
@@ -309,47 +337,6 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
raise ValueError("Unconditional is None but should not be")
|
||||
ip_hidden_states = torch.cat([self.adapter_ref().last_unconditional, ip_hidden_states], dim=0)
|
||||
|
||||
# `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 ip_hidden_states is not None:
|
||||
# apply scaler
|
||||
if self.train_scaler:
|
||||
@@ -365,8 +352,6 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
ip_value = ip_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
|
||||
ip_hidden_states = F.scaled_dot_product_attention(
|
||||
query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
@@ -376,26 +361,23 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
|
||||
scale = self.scale
|
||||
hidden_states = hidden_states + scale * ip_hidden_states
|
||||
# end ip adapter
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
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
|
||||
# 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)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
|
||||
# loosely based on # ref https://github.com/tencent-ailab/IP-Adapter/blob/main/tutorial_train.py
|
||||
class IPAdapter(torch.nn.Module):
|
||||
@@ -659,9 +641,9 @@ class IPAdapter(torch.nn.Module):
|
||||
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")
|
||||
# single transformer blocks do not have cross attn, but we will do them anyway
|
||||
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())
|
||||
|
||||
@@ -695,7 +677,7 @@ class IPAdapter(torch.nn.Module):
|
||||
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"):
|
||||
elif name.startswith("transformer") or name.startswith("single_transformer"):
|
||||
if is_flux:
|
||||
hidden_size = 3072
|
||||
else:
|
||||
@@ -773,11 +755,20 @@ class IPAdapter(torch.nn.Module):
|
||||
transformer: FluxTransformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"transformer_blocks.{i}.attn"]
|
||||
|
||||
# do single blocks too even though they dont have cross attn
|
||||
for i, module in transformer.single_transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"single_transformer_blocks.{i}.attn"]
|
||||
|
||||
self.adapter_modules = torch.nn.ModuleList(
|
||||
[
|
||||
transformer.transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.transformer_blocks))
|
||||
])
|
||||
] + [
|
||||
transformer.single_transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.single_transformer_blocks))
|
||||
]
|
||||
)
|
||||
else:
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
self.adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
@@ -1170,13 +1161,13 @@ class IPAdapter(torch.nn.Module):
|
||||
# when training just scaler, we do not train anything else
|
||||
if not self.config.train_scaler:
|
||||
param_groups.append({
|
||||
"params": self.get_non_scaler_parameters(),
|
||||
"params": list(self.get_non_scaler_parameters()),
|
||||
"lr": adapter_lr,
|
||||
})
|
||||
if self.config.train_scaler or self.config.merge_scaler:
|
||||
scaler_lr = adapter_lr if self.config.scaler_lr is None else self.config.scaler_lr
|
||||
param_groups.append({
|
||||
"params": self.get_scaler_parameters(),
|
||||
"params": list(self.get_scaler_parameters()),
|
||||
"lr": scaler_lr,
|
||||
})
|
||||
return param_groups
|
||||
|
||||
84
toolkit/logging.py
Normal file
84
toolkit/logging.py
Normal file
@@ -0,0 +1,84 @@
|
||||
from typing import OrderedDict, Optional
|
||||
from PIL import Image
|
||||
|
||||
from toolkit.config_modules import LoggingConfig
|
||||
|
||||
# Base logger class
|
||||
# This class does nothing, it's just a placeholder
|
||||
class EmptyLogger:
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
# start logging the training
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
# collect the log to send
|
||||
def log(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# send the log
|
||||
def commit(self, step: Optional[int] = None):
|
||||
pass
|
||||
|
||||
# log image
|
||||
def log_image(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# finish logging
|
||||
def finish(self):
|
||||
pass
|
||||
|
||||
# Wandb logger class
|
||||
# This class logs the data to wandb
|
||||
class WandbLogger(EmptyLogger):
|
||||
def __init__(self, project: str, run_name: str | None, config: OrderedDict) -> None:
|
||||
self.project = project
|
||||
self.run_name = run_name
|
||||
self.config = config
|
||||
|
||||
def start(self):
|
||||
try:
|
||||
import wandb
|
||||
except ImportError:
|
||||
raise ImportError("Failed to import wandb. Please install wandb by running `pip install wandb`")
|
||||
|
||||
# send the whole config to wandb
|
||||
run = wandb.init(project=self.project, name=self.run_name, config=self.config)
|
||||
self.run = run
|
||||
self._log = wandb.log # log function
|
||||
self._image = wandb.Image # image object
|
||||
|
||||
def log(self, *args, **kwargs):
|
||||
# when commit is False, wandb increments the step,
|
||||
# but we don't want that to happen, so we set commit=False
|
||||
self._log(*args, **kwargs, commit=False)
|
||||
|
||||
def commit(self, step: Optional[int] = None):
|
||||
# after overall one step is done, we commit the log
|
||||
# by log empty object with commit=True
|
||||
self._log({}, step=step, commit=True)
|
||||
|
||||
def log_image(
|
||||
self,
|
||||
image: Image,
|
||||
id, # sample index
|
||||
caption: str | None = None, # positive prompt
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
# create a wandb image object and log it
|
||||
image = self._image(image, caption=caption, *args, **kwargs)
|
||||
self._log({f"sample_{id}": image}, commit=False)
|
||||
|
||||
def finish(self):
|
||||
self.run.finish()
|
||||
|
||||
# create logger based on the logging config
|
||||
def create_logger(logging_config: LoggingConfig, all_config: OrderedDict):
|
||||
if logging_config.use_wandb:
|
||||
project_name = logging_config.project_name
|
||||
run_name = logging_config.run_name
|
||||
return WandbLogger(project=project_name, run_name=run_name, config=all_config)
|
||||
else:
|
||||
return EmptyLogger()
|
||||
@@ -9,6 +9,7 @@ from typing import List, Optional, Dict, Type, Union
|
||||
import torch
|
||||
from diffusers import UNet2DConditionModel, PixArtTransformer2DModel, AuraFlowTransformer2DModel
|
||||
from transformers import CLIPTextModel
|
||||
from toolkit.models.lokr import LokrModule
|
||||
|
||||
from .config_modules import NetworkConfig
|
||||
from .lorm import count_parameters
|
||||
@@ -19,9 +20,13 @@ sys.path.append(SD_SCRIPTS_ROOT)
|
||||
|
||||
from networks.lora import LoRANetwork, get_block_index
|
||||
from toolkit.models.DoRA import DoRAModule
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
|
||||
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
|
||||
|
||||
|
||||
@@ -63,7 +68,7 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
|
||||
torch.nn.Module.__init__(self)
|
||||
self.lora_name = lora_name
|
||||
self.orig_module_ref = weakref.ref(org_module)
|
||||
self.scalar = torch.tensor(1.0)
|
||||
self.scalar = torch.tensor(1.0, device=org_module.weight.device)
|
||||
# check if parent has bias. if not force use_bias to False
|
||||
if org_module.bias is None:
|
||||
use_bias = False
|
||||
@@ -163,6 +168,7 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
is_pixart: bool = False,
|
||||
is_auraflow: bool = False,
|
||||
is_flux: bool = False,
|
||||
is_lumina2: bool = False,
|
||||
use_bias: bool = False,
|
||||
is_lorm: bool = False,
|
||||
ignore_if_contains = None,
|
||||
@@ -176,6 +182,8 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
transformer_only: bool = False,
|
||||
peft_format: bool = False,
|
||||
is_assistant_adapter: bool = False,
|
||||
is_transformer: bool = False,
|
||||
base_model: 'StableDiffusion' = None,
|
||||
**kwargs
|
||||
) -> None:
|
||||
"""
|
||||
@@ -201,6 +209,9 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
ignore_if_contains = []
|
||||
self.ignore_if_contains = ignore_if_contains
|
||||
self.transformer_only = transformer_only
|
||||
self.base_model_ref = None
|
||||
if base_model is not None:
|
||||
self.base_model_ref = weakref.ref(base_model)
|
||||
|
||||
self.only_if_contains: Union[List, None] = only_if_contains
|
||||
|
||||
@@ -223,17 +234,26 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
self.is_pixart = is_pixart
|
||||
self.is_auraflow = is_auraflow
|
||||
self.is_flux = is_flux
|
||||
self.is_lumina2 = is_lumina2
|
||||
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
|
||||
elif self.network_type.lower() == "lokr":
|
||||
self.module_class = LokrModule
|
||||
module_class = LokrModule
|
||||
self.network_config: NetworkConfig = kwargs.get("network_config", None)
|
||||
|
||||
self.peft_format = peft_format
|
||||
self.is_transformer = is_transformer
|
||||
|
||||
|
||||
# always do peft for flux only for now
|
||||
if self.is_flux:
|
||||
self.peft_format = True
|
||||
if self.is_flux or self.is_v3 or self.is_lumina2 or is_transformer:
|
||||
# don't do peft format for lokr
|
||||
if self.network_type.lower() != "lokr":
|
||||
self.peft_format = True
|
||||
|
||||
if self.peft_format:
|
||||
# no alpha for peft
|
||||
@@ -273,7 +293,7 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
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:
|
||||
if is_pixart or is_v3 or is_auraflow or is_flux or is_lumina2 or self.is_transformer:
|
||||
unet_prefix = f"lora_transformer"
|
||||
if self.peft_format:
|
||||
unet_prefix = "transformer"
|
||||
@@ -305,15 +325,15 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
lora_name = ".".join(lora_name)
|
||||
# if it doesnt have a name, it wil have two dots
|
||||
lora_name.replace("..", ".")
|
||||
clean_name = lora_name
|
||||
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]):
|
||||
if any([word in clean_name for word in self.ignore_if_contains]):
|
||||
skip = True
|
||||
|
||||
# see if it is over threshold
|
||||
@@ -326,11 +346,27 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
if self.transformer_only and self.is_flux and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_lumina2 and is_unet:
|
||||
if "layers$$" not in lora_name and "noise_refiner$$" not in lora_name and "context_refiner$$" not in lora_name:
|
||||
skip = True
|
||||
if self.transformer_only and self.is_v3 and is_unet:
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
# handle custom models
|
||||
if self.transformer_only and is_unet and hasattr(root_module, 'transformer_blocks'):
|
||||
if "transformer_blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
if self.transformer_only and is_unet and hasattr(root_module, 'blocks'):
|
||||
if "blocks" not in lora_name:
|
||||
skip = True
|
||||
|
||||
if (is_linear or is_conv2d) and not skip:
|
||||
|
||||
if self.only_if_contains is not None and not any([word in lora_name for word in self.only_if_contains]):
|
||||
continue
|
||||
if self.only_if_contains is not None:
|
||||
if not any([word in clean_name for word in self.only_if_contains]) and not any([word in lora_name for word in self.only_if_contains]):
|
||||
continue
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
@@ -364,6 +400,11 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
self.conv_lora_dim is not None or conv_block_dims is not None):
|
||||
skipped.append(lora_name)
|
||||
continue
|
||||
|
||||
module_kwargs = {}
|
||||
|
||||
if self.network_type.lower() == "lokr":
|
||||
module_kwargs["factor"] = self.network_config.lokr_factor
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
@@ -377,10 +418,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
network=self,
|
||||
parent=module,
|
||||
use_bias=use_bias,
|
||||
**module_kwargs
|
||||
)
|
||||
loras.append(lora)
|
||||
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape), list(lora.lora_up.weight.shape)
|
||||
]
|
||||
if self.network_type.lower() == "lokr":
|
||||
try:
|
||||
lora_shape_dict[lora_name] = [list(lora.lokr_w1.weight.shape), list(lora.lokr_w2.weight.shape)]
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
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]
|
||||
@@ -428,6 +475,9 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
|
||||
|
||||
if is_flux:
|
||||
target_modules = ["FluxTransformer2DModel"]
|
||||
|
||||
if is_lumina2:
|
||||
target_modules = ["Lumina2Transformer2DModel"]
|
||||
|
||||
if train_unet:
|
||||
self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules)
|
||||
|
||||
@@ -354,7 +354,8 @@ def convert_diffusers_unet_to_lorm(
|
||||
elif child_module.__class__.__name__ in LINEAR_MODULES:
|
||||
if count_parameters(child_module) > parameter_threshold:
|
||||
|
||||
dtype = child_module.weight.dtype
|
||||
# dtype = child_module.weight.dtype
|
||||
dtype = torch.float32
|
||||
# extract and convert
|
||||
down_weight, up_weight, lora_dim, diff = extract_linear(
|
||||
weight=child_module.weight.clone().detach().float(),
|
||||
|
||||
1433
toolkit/models/base_model.py
Normal file
1433
toolkit/models/base_model.py
Normal file
File diff suppressed because it is too large
Load Diff
466
toolkit/models/cogview4.py
Normal file
466
toolkit/models/cogview4.py
Normal file
@@ -0,0 +1,466 @@
|
||||
# DONT USE THIS!. IT DOES NOT WORK YET!
|
||||
# Will revisit this when they release more info on how it was trained.
|
||||
|
||||
import weakref
|
||||
from diffusers import CogView4Pipeline
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from toolkit.basic import flush
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
|
||||
import os
|
||||
import copy
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, ModelArch
|
||||
import torch
|
||||
import diffusers
|
||||
from diffusers import AutoencoderKL, CogView4Transformer2DModel, CogView4Pipeline
|
||||
from optimum.quanto import freeze, qfloat8, QTensor, qint4
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from transformers import GlmModel, AutoTokenizer
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
|
||||
# remove this after a bug is fixed in diffusers code. This is a workaround.
|
||||
|
||||
|
||||
class FakeModel:
|
||||
def __init__(self, model):
|
||||
self.model_ref = weakref.ref(model)
|
||||
pass
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return self.model_ref().device
|
||||
|
||||
|
||||
scheduler_config = {
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.25,
|
||||
"invert_sigmas": False,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 0.75,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 1.0,
|
||||
"shift_terminal": None,
|
||||
"time_shift_type": "linear",
|
||||
"use_beta_sigmas": False,
|
||||
"use_dynamic_shifting": True,
|
||||
"use_exponential_sigmas": False,
|
||||
"use_karras_sigmas": False
|
||||
}
|
||||
|
||||
|
||||
class CogView4(BaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
model_config: ModelConfig,
|
||||
dtype='bf16',
|
||||
custom_pipeline=None,
|
||||
noise_scheduler=None,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(device, model_config, dtype,
|
||||
custom_pipeline, noise_scheduler, **kwargs)
|
||||
self.is_flow_matching = True
|
||||
self.is_transformer = True
|
||||
self.target_lora_modules = ['CogView4Transformer2DModel']
|
||||
|
||||
# cache for holding noise
|
||||
self.effective_noise = None
|
||||
|
||||
# static method to get the scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
scheduler = CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
return scheduler
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
base_model_path = "THUDM/CogView4-6B"
|
||||
model_path = self.model_config.name_or_path
|
||||
|
||||
self.print_and_status_update("Loading CogView4 model")
|
||||
# base_model_path = "black-forest-labs/FLUX.1-schnell"
|
||||
base_model_path = self.model_config.name_or_path_original
|
||||
subfolder = 'transformer'
|
||||
transformer_path = model_path
|
||||
if os.path.exists(transformer_path):
|
||||
subfolder = None
|
||||
transformer_path = os.path.join(transformer_path, 'transformer')
|
||||
# check if the path is a full checkpoint.
|
||||
te_folder_path = os.path.join(model_path, 'text_encoder')
|
||||
# if we have the te, this folder is a full checkpoint, use it as the base
|
||||
if os.path.exists(te_folder_path):
|
||||
base_model_path = model_path
|
||||
|
||||
self.print_and_status_update("Loading GlmModel")
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
base_model_path, subfolder="tokenizer", torch_dtype=dtype)
|
||||
text_encoder = GlmModel.from_pretrained(
|
||||
base_model_path, subfolder="text_encoder", torch_dtype=dtype)
|
||||
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing GlmModel")
|
||||
quantize(text_encoder, weights=get_qtype(self.model_config.qtype))
|
||||
freeze(text_encoder)
|
||||
flush()
|
||||
|
||||
# hack to fix diffusers bug workaround
|
||||
text_encoder.model = FakeModel(text_encoder)
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
transformer = CogView4Transformer2DModel.from_pretrained(
|
||||
transformer_path,
|
||||
subfolder=subfolder,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
|
||||
if self.model_config.split_model_over_gpus:
|
||||
raise ValueError(
|
||||
"Splitting model over gpus is not supported for CogViewModels models")
|
||||
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.assistant_lora_path is not None or self.model_config.inference_lora_path is not None:
|
||||
raise ValueError(
|
||||
"Assistant LoRA is not supported for CogViewModels models currently")
|
||||
|
||||
if self.model_config.lora_path is not None:
|
||||
raise ValueError(
|
||||
"Loading LoRA is not supported for CogViewModels models currently")
|
||||
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize:
|
||||
quantization_args = self.model_config.quantize_kwargs
|
||||
if 'exclude' not in quantization_args:
|
||||
quantization_args['exclude'] = []
|
||||
if 'include' not in quantization_args:
|
||||
quantization_args['include'] = []
|
||||
|
||||
# Be more specific with the include pattern to exactly match transformer blocks
|
||||
quantization_args['include'] += ["transformer_blocks.*"]
|
||||
|
||||
# Exclude all LayerNorm layers within transformer blocks
|
||||
quantization_args['exclude'] += [
|
||||
"transformer_blocks.*.norm1",
|
||||
"transformer_blocks.*.norm2",
|
||||
"transformer_blocks.*.norm2_context",
|
||||
"transformer_blocks.*.attn1.norm_q",
|
||||
"transformer_blocks.*.attn1.norm_k"
|
||||
]
|
||||
|
||||
# patch the state dict method
|
||||
patch_dequantization_on_save(transformer)
|
||||
quantization_type = get_qtype(self.model_config.qtype)
|
||||
self.print_and_status_update("Quantizing transformer")
|
||||
quantize(transformer, weights=quantization_type, **quantization_args)
|
||||
freeze(transformer)
|
||||
transformer.to(self.device_torch)
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
scheduler = CogView4.get_train_scheduler()
|
||||
self.print_and_status_update("Loading VAE")
|
||||
vae = AutoencoderKL.from_pretrained(
|
||||
base_model_path, subfolder="vae", torch_dtype=dtype)
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
pipe: CogView4Pipeline = CogView4Pipeline(
|
||||
scheduler=scheduler,
|
||||
text_encoder=None,
|
||||
tokenizer=tokenizer,
|
||||
vae=vae,
|
||||
transformer=None,
|
||||
)
|
||||
pipe.text_encoder = text_encoder
|
||||
pipe.transformer = transformer
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = pipe.text_encoder
|
||||
tokenizer = pipe.tokenizer
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
text_encoder.to(self.device_torch)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
self.pipeline = pipe
|
||||
self.model = transformer
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = CogView4.get_train_scheduler()
|
||||
pipeline = CogView4Pipeline(
|
||||
vae=self.vae,
|
||||
transformer=self.unet,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: CogView4Pipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
img = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
latents=gen_config.latents,
|
||||
generator=generator,
|
||||
**extra
|
||||
).images[0]
|
||||
return img
|
||||
|
||||
def get_noise_prediction(
|
||||
self,
|
||||
latent_model_input: torch.Tensor,
|
||||
timestep: torch.Tensor, # 0 to 1000 scale
|
||||
text_embeddings: PromptEmbeds,
|
||||
**kwargs
|
||||
):
|
||||
# target_size = (height, width)
|
||||
target_size = latent_model_input.shape[-2:]
|
||||
# multiply by 8
|
||||
target_size = (target_size[0] * 8, target_size[1] * 8)
|
||||
crops_coords_top_left = torch.tensor(
|
||||
[(0, 0)], dtype=self.torch_dtype, device=self.device_torch)
|
||||
|
||||
original_size = torch.tensor(
|
||||
[target_size], dtype=self.torch_dtype, device=self.device_torch)
|
||||
target_size = original_size.clone()
|
||||
noise_pred_cond = self.model(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
timestep=timestep,
|
||||
original_size=original_size,
|
||||
target_size=target_size,
|
||||
crop_coords=crops_coords_top_left,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
return noise_pred_cond
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
prompt_embeds, _ = self.pipeline.encode_prompt(
|
||||
prompt,
|
||||
do_classifier_free_guidance=False,
|
||||
device=self.device_torch,
|
||||
dtype=self.torch_dtype,
|
||||
)
|
||||
return PromptEmbeds(prompt_embeds)
|
||||
|
||||
def get_model_has_grad(self):
|
||||
return self.model.proj_out.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
return self.text_encoder.layers[0].mlp.down_proj.weight.requires_grad
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: CogView4Transformer2DModel = unwrap_model(self.model)
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, 'transformer'),
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
meta_path = os.path.join(output_path, 'aitk_meta.yaml')
|
||||
with open(meta_path, 'w') as f:
|
||||
yaml.dump(meta, f)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get('noise')
|
||||
effective_noise = self.effective_noise
|
||||
batch = kwargs.get('batch')
|
||||
if batch is None:
|
||||
raise ValueError("Batch is not provided")
|
||||
if noise is None:
|
||||
raise ValueError("Noise is not provided")
|
||||
# return batch.latents
|
||||
# return (batch.latents - noise).detach()
|
||||
return (noise - batch.latents).detach()
|
||||
# return (batch.latents).detach()
|
||||
# return (effective_noise - batch.latents).detach()
|
||||
|
||||
def _get_low_res_latents(self, latents):
|
||||
# todo prevent needing to do this and grab the tensor another way.
|
||||
with torch.no_grad():
|
||||
# Decode latents to image space
|
||||
images = self.decode_latents(
|
||||
latents, device=latents.device, dtype=latents.dtype)
|
||||
|
||||
# Downsample by a factor of 2 using bilinear interpolation
|
||||
B, C, H, W = images.shape
|
||||
low_res_images = torch.nn.functional.interpolate(
|
||||
images,
|
||||
size=(H // 2, W // 2),
|
||||
mode="bilinear",
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# Upsample back to original resolution to match expected VAE input dimensions
|
||||
upsampled_low_res_images = torch.nn.functional.interpolate(
|
||||
low_res_images,
|
||||
size=(H, W),
|
||||
mode="bilinear",
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
# Encode the low-resolution images back to latent space
|
||||
low_res_latents = self.encode_images(
|
||||
upsampled_low_res_images, device=latents.device, dtype=latents.dtype)
|
||||
return low_res_latents
|
||||
|
||||
# def add_noise(
|
||||
# self,
|
||||
# original_samples: torch.FloatTensor,
|
||||
# noise: torch.FloatTensor,
|
||||
# timesteps: torch.IntTensor,
|
||||
# **kwargs,
|
||||
# ) -> torch.FloatTensor:
|
||||
# relay_start_point = 500
|
||||
|
||||
# # Store original samples for loss calculation
|
||||
# self.original_samples = original_samples
|
||||
|
||||
# # Prepare chunks for batch processing
|
||||
# original_samples_chunks = torch.chunk(
|
||||
# original_samples, original_samples.shape[0], dim=0)
|
||||
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
|
||||
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
|
||||
|
||||
# # Get the low res latents only if needed
|
||||
# low_res_latents_chunks = None
|
||||
|
||||
# # Handle case where timesteps is a single value for all samples
|
||||
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
|
||||
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
|
||||
|
||||
# noisy_latents_chunks = []
|
||||
# effective_noise_chunks = [] # Store the effective noise for each sample
|
||||
|
||||
# for idx in range(original_samples.shape[0]):
|
||||
# t = timesteps_chunks[idx]
|
||||
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
|
||||
|
||||
# # Flowmatching interpolation between original and noise
|
||||
# if t > relay_start_point:
|
||||
# # Standard flowmatching - direct linear interpolation
|
||||
# noisy_latents = (1 - t_01) * original_samples_chunks[idx] + t_01 * noise_chunks[idx]
|
||||
# effective_noise_chunks.append(noise_chunks[idx]) # Effective noise is just the noise
|
||||
# else:
|
||||
# # Relay flowmatching case - only compute low_res_latents if needed
|
||||
# if low_res_latents_chunks is None:
|
||||
# low_res_latents = self._get_low_res_latents(original_samples)
|
||||
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
|
||||
|
||||
# # Calculate the relay ratio (0 to 1)
|
||||
# t_ratio = t.float() / relay_start_point
|
||||
# t_ratio = torch.clamp(t_ratio, 0.0, 1.0)
|
||||
|
||||
# # First blend between original and low-res based on t_ratio
|
||||
# z0_t = (1 - t_ratio) * original_samples_chunks[idx] + t_ratio * low_res_latents_chunks[idx]
|
||||
|
||||
# added_lor_res_noise = z0_t - original_samples_chunks[idx]
|
||||
|
||||
# # Then apply flowmatching interpolation between this blended state and noise
|
||||
# noisy_latents = (1 - t_01) * z0_t + t_01 * noise_chunks[idx]
|
||||
|
||||
# # For prediction target, we need to store the effective "source"
|
||||
# effective_noise_chunks.append(noise_chunks[idx] + added_lor_res_noise)
|
||||
|
||||
# noisy_latents_chunks.append(noisy_latents)
|
||||
|
||||
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
|
||||
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
|
||||
|
||||
# return noisy_latents
|
||||
|
||||
# def add_noise(
|
||||
# self,
|
||||
# original_samples: torch.FloatTensor,
|
||||
# noise: torch.FloatTensor,
|
||||
# timesteps: torch.IntTensor,
|
||||
# **kwargs,
|
||||
# ) -> torch.FloatTensor:
|
||||
# relay_start_point = 500
|
||||
|
||||
# # Store original samples for loss calculation
|
||||
# self.original_samples = original_samples
|
||||
|
||||
# # Prepare chunks for batch processing
|
||||
# original_samples_chunks = torch.chunk(
|
||||
# original_samples, original_samples.shape[0], dim=0)
|
||||
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
|
||||
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
|
||||
|
||||
# # Get the low res latents only if needed
|
||||
# low_res_latents = self._get_low_res_latents(original_samples)
|
||||
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
|
||||
|
||||
# # Handle case where timesteps is a single value for all samples
|
||||
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
|
||||
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
|
||||
|
||||
# noisy_latents_chunks = []
|
||||
# effective_noise_chunks = [] # Store the effective noise for each sample
|
||||
|
||||
# for idx in range(original_samples.shape[0]):
|
||||
# t = timesteps_chunks[idx]
|
||||
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
|
||||
|
||||
# lrln = low_res_latents_chunks[idx] - original_samples_chunks[idx]
|
||||
# # lrln = lrln * (1 - t_01)
|
||||
|
||||
# # make the noise an interpolation between noise and low_res_latents with
|
||||
# # being noise at t_01=1 and low_res_latents at t_01=0
|
||||
# new_noise = t_01 * noise_chunks[idx] + (1 - t_01) * lrln
|
||||
# # new_noise = noise_chunks[idx] + lrln
|
||||
# # new_noise = noise_chunks[idx] + lrln
|
||||
|
||||
# # Then apply flowmatching interpolation between this blended state and noise
|
||||
# noisy_latents = (1 - t_01) * original_samples + t_01 * new_noise
|
||||
|
||||
# # For prediction target, we need to store the effective "source"
|
||||
# effective_noise_chunks.append(new_noise)
|
||||
|
||||
# noisy_latents_chunks.append(noisy_latents)
|
||||
|
||||
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
|
||||
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
|
||||
|
||||
# return noisy_latents
|
||||
272
toolkit/models/control_lora_adapter.py
Normal file
272
toolkit/models/control_lora_adapter.py
Normal file
@@ -0,0 +1,272 @@
|
||||
import inspect
|
||||
import weakref
|
||||
import torch
|
||||
from typing import TYPE_CHECKING
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
from diffusers import FluxTransformer2DModel
|
||||
# weakref
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.config_modules import AdapterConfig, TrainConfig, ModelConfig
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
# after each step we concat the control image with the latents
|
||||
# latent_model_input = torch.cat([latents, control_image], dim=2)
|
||||
# the x_embedder has a full rank lora to handle the additional channels
|
||||
# this replaces the x_embedder with a full rank lora. on flux this is
|
||||
# x_embedder(diffusers) or img_in(bfl)
|
||||
|
||||
# Flux
|
||||
# img_in.lora_A.weight [128, 128]
|
||||
# img_in.lora_B.bias [3 072]
|
||||
# img_in.lora_B.weight [3 072, 128]
|
||||
|
||||
|
||||
class ImgEmbedder(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'ControlLoraAdapter',
|
||||
orig_layer: torch.nn.Linear,
|
||||
in_channels=64,
|
||||
out_channels=3072
|
||||
):
|
||||
super().__init__()
|
||||
# only do the weight for the new input. We combine with the original linear layer
|
||||
init = torch.randn(out_channels, in_channels, device=orig_layer.weight.device, dtype=orig_layer.weight.dtype) * 0.01
|
||||
self.weight = torch.nn.Parameter(init)
|
||||
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.orig_layer_ref: weakref.ref = weakref.ref(orig_layer)
|
||||
|
||||
@classmethod
|
||||
def from_model(
|
||||
cls,
|
||||
model: FluxTransformer2DModel,
|
||||
adapter: 'ControlLoraAdapter',
|
||||
num_control_images=1,
|
||||
has_inpainting_input=False
|
||||
):
|
||||
if model.__class__.__name__ == 'FluxTransformer2DModel':
|
||||
num_adapter_in_channels = model.x_embedder.in_features * num_control_images
|
||||
|
||||
if has_inpainting_input:
|
||||
# inpainting has the mask before packing latents. it is normally 16 ch + 1ch mask
|
||||
# packed it is 64ch + 4ch mask
|
||||
# so we need to add 4 to the input channels
|
||||
num_adapter_in_channels += 4
|
||||
|
||||
x_embedder: torch.nn.Linear = model.x_embedder
|
||||
img_embedder = cls(
|
||||
adapter,
|
||||
orig_layer=x_embedder,
|
||||
in_channels=num_adapter_in_channels,
|
||||
out_channels=x_embedder.out_features,
|
||||
)
|
||||
|
||||
# hijack the forward method
|
||||
x_embedder._orig_ctrl_lora_forward = x_embedder.forward
|
||||
x_embedder.forward = img_embedder.forward
|
||||
|
||||
# update the config of the transformer
|
||||
model.config.in_channels = model.config.in_channels * (num_control_images + 1)
|
||||
model.config["in_channels"] = model.config.in_channels
|
||||
|
||||
return img_embedder
|
||||
else:
|
||||
raise ValueError("Model not supported")
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
if not self.is_active:
|
||||
# make sure lora is not active
|
||||
if self.adapter_ref().control_lora is not None:
|
||||
self.adapter_ref().control_lora.is_active = False
|
||||
return self.orig_layer_ref()._orig_ctrl_lora_forward(x)
|
||||
|
||||
# make sure lora is active
|
||||
if self.adapter_ref().control_lora is not None:
|
||||
self.adapter_ref().control_lora.is_active = True
|
||||
|
||||
orig_device = x.device
|
||||
orig_dtype = x.dtype
|
||||
|
||||
x = x.to(self.weight.device, dtype=self.weight.dtype)
|
||||
|
||||
orig_weight = self.orig_layer_ref().weight.data.detach()
|
||||
orig_weight = orig_weight.to(self.weight.device, dtype=self.weight.dtype)
|
||||
linear_weight = torch.cat([orig_weight, self.weight], dim=1)
|
||||
|
||||
bias = None
|
||||
if self.orig_layer_ref().bias is not None:
|
||||
bias = self.orig_layer_ref().bias.data.detach().to(self.weight.device, dtype=self.weight.dtype)
|
||||
|
||||
x = torch.nn.functional.linear(x, linear_weight, bias)
|
||||
|
||||
x = x.to(orig_device, dtype=orig_dtype)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
class ControlLoraAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
config: 'AdapterConfig',
|
||||
train_config: 'TrainConfig'
|
||||
):
|
||||
super().__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref = weakref.ref(sd)
|
||||
self.model_config: ModelConfig = sd.model_config
|
||||
self.network_config = config.lora_config
|
||||
self.train_config = train_config
|
||||
self.device_torch = sd.device_torch
|
||||
self.control_lora = None
|
||||
|
||||
if self.network_config is not None:
|
||||
|
||||
network_kwargs = {} if self.network_config.network_kwargs is None else self.network_config.network_kwargs
|
||||
if hasattr(sd, 'target_lora_modules'):
|
||||
network_kwargs['target_lin_modules'] = self.sd.target_lora_modules
|
||||
|
||||
if 'ignore_if_contains' not in network_kwargs:
|
||||
network_kwargs['ignore_if_contains'] = []
|
||||
|
||||
# always ignore x_embedder
|
||||
network_kwargs['ignore_if_contains'].append('x_embedder')
|
||||
|
||||
self.control_lora = LoRASpecialNetwork(
|
||||
text_encoder=sd.text_encoder,
|
||||
unet=sd.unet,
|
||||
lora_dim=self.network_config.linear,
|
||||
multiplier=1.0,
|
||||
alpha=self.network_config.linear_alpha,
|
||||
train_unet=self.train_config.train_unet,
|
||||
train_text_encoder=self.train_config.train_text_encoder,
|
||||
conv_lora_dim=self.network_config.conv,
|
||||
conv_alpha=self.network_config.conv_alpha,
|
||||
is_sdxl=self.model_config.is_xl or self.model_config.is_ssd,
|
||||
is_v2=self.model_config.is_v2,
|
||||
is_v3=self.model_config.is_v3,
|
||||
is_pixart=self.model_config.is_pixart,
|
||||
is_auraflow=self.model_config.is_auraflow,
|
||||
is_flux=self.model_config.is_flux,
|
||||
is_lumina2=self.model_config.is_lumina2,
|
||||
is_ssd=self.model_config.is_ssd,
|
||||
is_vega=self.model_config.is_vega,
|
||||
dropout=self.network_config.dropout,
|
||||
use_text_encoder_1=self.model_config.use_text_encoder_1,
|
||||
use_text_encoder_2=self.model_config.use_text_encoder_2,
|
||||
use_bias=False,
|
||||
is_lorm=False,
|
||||
network_config=self.network_config,
|
||||
network_type=self.network_config.type,
|
||||
transformer_only=self.network_config.transformer_only,
|
||||
is_transformer=sd.is_transformer,
|
||||
base_model=sd,
|
||||
**network_kwargs
|
||||
)
|
||||
self.control_lora.force_to(self.device_torch, dtype=torch.float32)
|
||||
self.control_lora._update_torch_multiplier()
|
||||
self.control_lora.apply_to(
|
||||
sd.text_encoder,
|
||||
sd.unet,
|
||||
self.train_config.train_text_encoder,
|
||||
self.train_config.train_unet
|
||||
)
|
||||
self.control_lora.can_merge_in = False
|
||||
self.control_lora.prepare_grad_etc(sd.text_encoder, sd.unet)
|
||||
if self.train_config.gradient_checkpointing:
|
||||
self.control_lora.enable_gradient_checkpointing()
|
||||
|
||||
self.x_embedder = ImgEmbedder.from_model(
|
||||
sd.unet,
|
||||
self,
|
||||
num_control_images=config.num_control_images,
|
||||
has_inpainting_input=config.has_inpainting_input
|
||||
)
|
||||
self.x_embedder.to(self.device_torch)
|
||||
|
||||
def get_params(self):
|
||||
if self.control_lora is not None:
|
||||
config = {
|
||||
'text_encoder_lr': self.train_config.lr,
|
||||
'unet_lr': self.train_config.lr,
|
||||
}
|
||||
sig = inspect.signature(self.control_lora.prepare_optimizer_params)
|
||||
if 'default_lr' in sig.parameters:
|
||||
config['default_lr'] = self.train_config.lr
|
||||
if 'learning_rate' in sig.parameters:
|
||||
config['learning_rate'] = self.train_config.lr
|
||||
params_net = self.control_lora.prepare_optimizer_params(
|
||||
**config
|
||||
)
|
||||
|
||||
# we want only tensors here
|
||||
params = []
|
||||
for p in params_net:
|
||||
if isinstance(p, dict):
|
||||
params += p["params"]
|
||||
elif isinstance(p, torch.Tensor):
|
||||
params.append(p)
|
||||
elif isinstance(p, list):
|
||||
params += p
|
||||
else:
|
||||
params = []
|
||||
|
||||
# make sure the embedder is float32
|
||||
self.x_embedder.to(torch.float32)
|
||||
|
||||
params += list(self.x_embedder.parameters())
|
||||
|
||||
# we need to be able to yield from the list like yield from params
|
||||
|
||||
return params
|
||||
|
||||
def load_weights(self, state_dict, strict=True):
|
||||
lora_sd = {}
|
||||
img_embedder_sd = {}
|
||||
for key, value in state_dict.items():
|
||||
if "x_embedder" in key:
|
||||
new_key = key.replace("transformer.x_embedder.", "")
|
||||
img_embedder_sd[new_key] = value
|
||||
else:
|
||||
lora_sd[key] = value
|
||||
|
||||
# todo process state dict before loading
|
||||
if self.control_lora is not None:
|
||||
self.control_lora.load_weights(lora_sd)
|
||||
# automatically upgrade the x imbedder if more dims are added
|
||||
if self.x_embedder.weight.shape[1] > img_embedder_sd['weight'].shape[1]:
|
||||
print("Upgrading x_embedder from {} to {}".format(
|
||||
img_embedder_sd['weight'].shape[1],
|
||||
self.x_embedder.weight.shape[1]
|
||||
))
|
||||
while img_embedder_sd['weight'].shape[1] < self.x_embedder.weight.shape[1]:
|
||||
img_embedder_sd['weight'] = torch.cat([img_embedder_sd['weight'] ] * 2, dim=1)
|
||||
if img_embedder_sd['weight'].shape[1] > self.x_embedder.weight.shape[1]:
|
||||
img_embedder_sd['weight'] = img_embedder_sd['weight'][:, :self.x_embedder.weight.shape[1]]
|
||||
self.x_embedder.load_state_dict(img_embedder_sd, strict=False)
|
||||
|
||||
def get_state_dict(self):
|
||||
if self.control_lora is not None:
|
||||
lora_sd = self.control_lora.get_state_dict(dtype=torch.float32)
|
||||
else:
|
||||
lora_sd = {}
|
||||
# todo make sure we match loras elseware.
|
||||
img_embedder_sd = self.x_embedder.state_dict()
|
||||
for key, value in img_embedder_sd.items():
|
||||
lora_sd[f"transformer.x_embedder.{key}"] = value
|
||||
return lora_sd
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
33
toolkit/models/decorator.py
Normal file
33
toolkit/models/decorator.py
Normal file
@@ -0,0 +1,33 @@
|
||||
import torch
|
||||
|
||||
|
||||
class Decorator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int = 4,
|
||||
token_size: int = 4096,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.weight: torch.nn.Parameter = torch.nn.Parameter(
|
||||
torch.randn(num_tokens, token_size)
|
||||
)
|
||||
# ensure it is float32
|
||||
self.weight.data = self.weight.data.float()
|
||||
|
||||
def forward(self, text_embeds: torch.Tensor, is_unconditional=False) -> torch.Tensor:
|
||||
# make sure the param is float32
|
||||
if self.weight.dtype != text_embeds.dtype:
|
||||
self.weight.data = self.weight.data.float()
|
||||
# expand batch to match text_embeds
|
||||
batch_size = text_embeds.shape[0]
|
||||
decorator_embeds = self.weight.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
if is_unconditional:
|
||||
# zero pad the decorator embeds
|
||||
decorator_embeds = torch.zeros_like(decorator_embeds)
|
||||
|
||||
if decorator_embeds.dtype != text_embeds.dtype:
|
||||
decorator_embeds = decorator_embeds.to(text_embeds.dtype)
|
||||
text_embeds = torch.cat((text_embeds, decorator_embeds), dim=-2)
|
||||
|
||||
return text_embeds
|
||||
367
toolkit/models/diffusion_feature_extraction.py
Normal file
367
toolkit/models/diffusion_feature_extraction.py
Normal file
@@ -0,0 +1,367 @@
|
||||
import torch
|
||||
import os
|
||||
from torch import nn
|
||||
from safetensors.torch import load_file
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderTiny
|
||||
from transformers import SiglipImageProcessor, SiglipVisionModel
|
||||
import lpips
|
||||
|
||||
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
|
||||
|
||||
class ResBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
|
||||
self.norm1 = nn.GroupNorm(8, out_channels)
|
||||
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
|
||||
self.norm2 = nn.GroupNorm(8, out_channels)
|
||||
self.skip = nn.Conv2d(in_channels, out_channels,
|
||||
1) if in_channels != out_channels else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
identity = self.skip(x)
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = F.silu(x)
|
||||
x = self.conv2(x)
|
||||
x = self.norm2(x)
|
||||
x = F.silu(x + identity)
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor2(nn.Module):
|
||||
def __init__(self, in_channels=32):
|
||||
super().__init__()
|
||||
self.version = 2
|
||||
|
||||
# Path 1: Upsample to 512x512 (1, 64, 512, 512)
|
||||
self.up_path = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 64, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(64, 64),
|
||||
nn.Conv2d(64, 64, 3, padding=1),
|
||||
])
|
||||
|
||||
# Path 2: Upsample to 256x256 (1, 128, 256, 256)
|
||||
self.path2 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 128, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(128, 128),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(128, 128),
|
||||
nn.Conv2d(128, 128, 3, padding=1),
|
||||
])
|
||||
|
||||
# Path 3: Upsample to 128x128 (1, 256, 128, 128)
|
||||
self.path3 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 256, 3, padding=1),
|
||||
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
|
||||
ResBlock(256, 256),
|
||||
nn.Conv2d(256, 256, 3, padding=1)
|
||||
])
|
||||
|
||||
# Path 4: Original size (1, 512, 64, 64)
|
||||
self.path4 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 512, 3, padding=1),
|
||||
ResBlock(512, 512),
|
||||
ResBlock(512, 512),
|
||||
nn.Conv2d(512, 512, 3, padding=1)
|
||||
])
|
||||
|
||||
# Path 5: Downsample to 32x32 (1, 512, 32, 32)
|
||||
self.path5 = nn.ModuleList([
|
||||
nn.Conv2d(in_channels, 512, 3, padding=1),
|
||||
ResBlock(512, 512),
|
||||
nn.AvgPool2d(2),
|
||||
ResBlock(512, 512),
|
||||
nn.Conv2d(512, 512, 3, padding=1)
|
||||
])
|
||||
|
||||
def forward(self, x):
|
||||
outputs = []
|
||||
|
||||
# Path 1: 512x512
|
||||
x1 = x
|
||||
for layer in self.up_path:
|
||||
x1 = layer(x1)
|
||||
outputs.append(x1) # [1, 64, 512, 512]
|
||||
|
||||
# Path 2: 256x256
|
||||
x2 = x
|
||||
for layer in self.path2:
|
||||
x2 = layer(x2)
|
||||
outputs.append(x2) # [1, 128, 256, 256]
|
||||
|
||||
# Path 3: 128x128
|
||||
x3 = x
|
||||
for layer in self.path3:
|
||||
x3 = layer(x3)
|
||||
outputs.append(x3) # [1, 256, 128, 128]
|
||||
|
||||
# Path 4: 64x64
|
||||
x4 = x
|
||||
for layer in self.path4:
|
||||
x4 = layer(x4)
|
||||
outputs.append(x4) # [1, 512, 64, 64]
|
||||
|
||||
# Path 5: 32x32
|
||||
x5 = x
|
||||
for layer in self.path5:
|
||||
x5 = layer(x5)
|
||||
outputs.append(x5) # [1, 512, 32, 32]
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
class DFEBlock(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)
|
||||
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)
|
||||
self.act = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
x_in = x
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x)
|
||||
x = self.act(x)
|
||||
x = x + x_in
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor(nn.Module):
|
||||
def __init__(self, in_channels=32):
|
||||
super().__init__()
|
||||
self.version = 1
|
||||
num_blocks = 6
|
||||
self.conv_in = nn.Conv2d(in_channels, 512, 1)
|
||||
self.blocks = nn.ModuleList([DFEBlock(512) for _ in range(num_blocks)])
|
||||
self.conv_out = nn.Conv2d(512, 512, 1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
for block in self.blocks:
|
||||
x = block(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class DiffusionFeatureExtractor3(nn.Module):
|
||||
def __init__(self, device=torch.device("cuda"), dtype=torch.bfloat16):
|
||||
super().__init__()
|
||||
self.version = 3
|
||||
vae = AutoencoderTiny.from_pretrained(
|
||||
"madebyollin/taef1", torch_dtype=torch.bfloat16)
|
||||
self.vae = vae
|
||||
image_encoder_path = "google/siglip-so400m-patch14-384"
|
||||
try:
|
||||
self.image_processor = SiglipImageProcessor.from_pretrained(
|
||||
image_encoder_path)
|
||||
except EnvironmentError:
|
||||
self.image_processor = SiglipImageProcessor()
|
||||
self.vision_encoder = SiglipVisionModel.from_pretrained(
|
||||
image_encoder_path,
|
||||
ignore_mismatched_sizes=True
|
||||
).to(device, dtype=dtype)
|
||||
|
||||
self.lpips_model = lpips_model = lpips.LPIPS(net='vgg')
|
||||
self.lpips_model = lpips_model.to(device, dtype=torch.float32)
|
||||
self.losses = {}
|
||||
self.log_every = 100
|
||||
self.step = 0
|
||||
|
||||
def get_siglip_features(self, tensors_0_1):
|
||||
dtype = torch.bfloat16
|
||||
device = self.vae.device
|
||||
# resize to 384x384
|
||||
images = F.interpolate(tensors_0_1, size=(384, 384),
|
||||
mode='bicubic', align_corners=False)
|
||||
|
||||
mean = torch.tensor(self.image_processor.image_mean).to(
|
||||
device, dtype=dtype
|
||||
).detach()
|
||||
std = torch.tensor(self.image_processor.image_std).to(
|
||||
device, dtype=dtype
|
||||
).detach()
|
||||
# tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
|
||||
clip_image = (
|
||||
images - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
|
||||
id_embeds = self.vision_encoder(
|
||||
clip_image,
|
||||
output_hidden_states=True,
|
||||
)
|
||||
|
||||
last_hidden_state = id_embeds['last_hidden_state']
|
||||
return last_hidden_state
|
||||
|
||||
def get_lpips_features(self, tensors_0_1):
|
||||
device = self.vae.device
|
||||
tensors_n1p1 = (tensors_0_1 * 2) - 1
|
||||
def get_lpips_features(img): # -1 to 1
|
||||
in0_input = self.lpips_model.scaling_layer(img)
|
||||
outs0 = self.lpips_model.net.forward(in0_input)
|
||||
|
||||
feats0 = {}
|
||||
|
||||
feats_list = []
|
||||
for kk in range(self.lpips_model.L):
|
||||
feats0[kk] = lpips.normalize_tensor(outs0[kk])
|
||||
feats_list.append(feats0[kk])
|
||||
|
||||
# 512 in
|
||||
# vgg
|
||||
# 0 torch.Size([1, 64, 512, 512])
|
||||
# 1 torch.Size([1, 128, 256, 256])
|
||||
# 2 torch.Size([1, 256, 128, 128])
|
||||
# 3 torch.Size([1, 512, 64, 64])
|
||||
# 4 torch.Size([1, 512, 32, 32])
|
||||
|
||||
return feats_list
|
||||
|
||||
# do lpips
|
||||
lpips_feat_list = [x for x in get_lpips_features(
|
||||
tensors_n1p1.to(device, dtype=torch.float32))]
|
||||
|
||||
return lpips_feat_list
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
noise,
|
||||
noise_pred,
|
||||
noisy_latents,
|
||||
timesteps,
|
||||
batch: DataLoaderBatchDTO,
|
||||
scheduler: CustomFlowMatchEulerDiscreteScheduler,
|
||||
# lpips_weight=1.0,
|
||||
lpips_weight=10.0,
|
||||
clip_weight=0.1,
|
||||
pixel_weight=0.1
|
||||
):
|
||||
dtype = torch.bfloat16
|
||||
device = self.vae.device
|
||||
|
||||
# first we step the scheduler from current timestep to the very end for a full denoise
|
||||
# bs = noise_pred.shape[0]
|
||||
# noise_pred_chunks = torch.chunk(noise_pred, bs)
|
||||
# timestep_chunks = torch.chunk(timesteps, bs)
|
||||
# noisy_latent_chunks = torch.chunk(noisy_latents, bs)
|
||||
# stepped_chunks = []
|
||||
# for idx in range(bs):
|
||||
# model_output = noise_pred_chunks[idx]
|
||||
# timestep = timestep_chunks[idx]
|
||||
# scheduler._step_index = None
|
||||
# scheduler._init_step_index(timestep)
|
||||
# sample = noisy_latent_chunks[idx].to(torch.float32)
|
||||
|
||||
# sigma = scheduler.sigmas[scheduler.step_index]
|
||||
# sigma_next = scheduler.sigmas[-1] # use last sigma for final step
|
||||
# prev_sample = sample + (sigma_next - sigma) * model_output
|
||||
# stepped_chunks.append(prev_sample)
|
||||
|
||||
# stepped_latents = torch.cat(stepped_chunks, dim=0)
|
||||
|
||||
stepped_latents = noise - noise_pred
|
||||
|
||||
latents = stepped_latents.to(self.vae.device, dtype=self.vae.dtype)
|
||||
|
||||
latents = (
|
||||
latents / self.vae.config['scaling_factor']) + self.vae.config['shift_factor']
|
||||
tensors_n1p1 = self.vae.decode(latents).sample # -1 to 1
|
||||
|
||||
pred_images = (tensors_n1p1 + 1) / 2 # 0 to 1
|
||||
|
||||
lpips_feat_list_pred = self.get_lpips_features(pred_images.float())
|
||||
|
||||
total_loss = 0
|
||||
|
||||
with torch.no_grad():
|
||||
target_img = batch.tensor.to(device, dtype=dtype)
|
||||
# go from -1 to 1 to 0 to 1
|
||||
target_img = (target_img + 1) / 2
|
||||
lpips_feat_list_target = self.get_lpips_features(target_img.float())
|
||||
if clip_weight > 0:
|
||||
target_clip_output = self.get_siglip_features(target_img).detach()
|
||||
if clip_weight > 0:
|
||||
pred_clip_output = self.get_siglip_features(pred_images)
|
||||
clip_loss = torch.nn.functional.mse_loss(
|
||||
pred_clip_output.float(), target_clip_output.float()
|
||||
) * clip_weight
|
||||
|
||||
if 'clip_loss' not in self.losses:
|
||||
self.losses['clip_loss'] = clip_loss.item()
|
||||
else:
|
||||
self.losses['clip_loss'] += clip_loss.item()
|
||||
|
||||
total_loss += clip_loss
|
||||
|
||||
skip_lpips_layers = []
|
||||
|
||||
lpips_loss = 0
|
||||
for idx, lpips_feat in enumerate(lpips_feat_list_pred):
|
||||
if idx in skip_lpips_layers:
|
||||
continue
|
||||
lpips_loss += torch.nn.functional.mse_loss(
|
||||
lpips_feat.float(), lpips_feat_list_target[idx].float()
|
||||
) * lpips_weight
|
||||
|
||||
if f'lpips_loss_{idx}' not in self.losses:
|
||||
self.losses[f'lpips_loss_{idx}'] = lpips_loss.item()
|
||||
else:
|
||||
self.losses[f'lpips_loss_{idx}'] += lpips_loss.item()
|
||||
|
||||
total_loss += lpips_loss
|
||||
|
||||
# mse_loss = torch.nn.functional.mse_loss(
|
||||
# stepped_latents.float(), batch.latents.float()
|
||||
# ) * pixel_weight
|
||||
|
||||
# if 'pixel_loss' not in self.losses:
|
||||
# self.losses['pixel_loss'] = mse_loss.item()
|
||||
# else:
|
||||
# self.losses['pixel_loss'] += mse_loss.item()
|
||||
|
||||
if self.step % self.log_every == 0 and self.step > 0:
|
||||
print(f"DFE losses:")
|
||||
for key in self.losses:
|
||||
self.losses[key] /= self.log_every
|
||||
# print in 2.000e-01 format
|
||||
print(f" - {key}: {self.losses[key]:.3e}")
|
||||
self.losses[key] = 0.0
|
||||
|
||||
# total_loss += mse_loss
|
||||
self.step += 1
|
||||
|
||||
return total_loss
|
||||
|
||||
|
||||
def load_dfe(model_path) -> DiffusionFeatureExtractor:
|
||||
if model_path == "v3":
|
||||
dfe = DiffusionFeatureExtractor3()
|
||||
dfe.eval()
|
||||
return dfe
|
||||
if not os.path.exists(model_path):
|
||||
raise FileNotFoundError(f"Model file not found: {model_path}")
|
||||
# if it ende with safetensors
|
||||
if model_path.endswith('.safetensors'):
|
||||
state_dict = load_file(model_path)
|
||||
else:
|
||||
state_dict = torch.load(model_path, weights_only=True)
|
||||
if 'model_state_dict' in state_dict:
|
||||
state_dict = state_dict['model_state_dict']
|
||||
|
||||
if 'conv_in.weight' in state_dict:
|
||||
dfe = DiffusionFeatureExtractor()
|
||||
else:
|
||||
dfe = DiffusionFeatureExtractor2()
|
||||
|
||||
dfe.load_state_dict(state_dict)
|
||||
dfe.eval()
|
||||
return dfe
|
||||
993
toolkit/models/flex2.py
Normal file
993
toolkit/models/flex2.py
Normal file
@@ -0,0 +1,993 @@
|
||||
from typing import List, Optional, Union
|
||||
from diffusers import FluxPipeline
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.loaders import FluxLoraLoaderMixin, TextualInversionLoaderMixin
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
is_torch_xla_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
scale_lora_layers,
|
||||
unscale_lora_layers,
|
||||
)
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
from transformers import (
|
||||
CLIPImageProcessor,
|
||||
CLIPTextModel,
|
||||
CLIPTokenizer,
|
||||
CLIPVisionModelWithProjection
|
||||
)
|
||||
|
||||
|
||||
|
||||
from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
|
||||
from diffusers.loaders import FluxIPAdapterMixin, FluxLoraLoaderMixin, FromSingleFileMixin, TextualInversionLoaderMixin
|
||||
from diffusers.models import AutoencoderKL, FluxTransformer2DModel
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import Flex2Pipeline
|
||||
|
||||
>>> pipe = Flex2Pipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16)
|
||||
>>> pipe.to("cuda")
|
||||
>>> prompt = "A cat holding a sign that says hello world"
|
||||
>>> # Depending on the variant being used, the pipeline call will slightly vary.
|
||||
>>> # Refer to the pipeline documentation for more details.
|
||||
>>> image = pipe(prompt, num_inference_steps=4, guidance_scale=0.0).images[0]
|
||||
>>> image.save("flux.png")
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len,
|
||||
base_seq_len: int = 256,
|
||||
max_seq_len: int = 4096,
|
||||
base_shift: float = 0.5,
|
||||
max_shift: float = 1.16,
|
||||
):
|
||||
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
b = base_shift - m * base_seq_len
|
||||
mu = image_seq_len * m + b
|
||||
return mu
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
|
||||
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
|
||||
|
||||
Args:
|
||||
scheduler (`SchedulerMixin`):
|
||||
The scheduler to get timesteps from.
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
|
||||
must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
|
||||
`num_inference_steps` and `sigmas` must be `None`.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
|
||||
`num_inference_steps` and `timesteps` must be `None`.
|
||||
|
||||
Returns:
|
||||
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
|
||||
second element is the number of inference steps.
|
||||
"""
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class Flex2Pipeline(
|
||||
DiffusionPipeline,
|
||||
FluxLoraLoaderMixin,
|
||||
FromSingleFileMixin,
|
||||
TextualInversionLoaderMixin,
|
||||
FluxIPAdapterMixin,
|
||||
):
|
||||
r"""
|
||||
The Flux pipeline for text-to-image generation.
|
||||
|
||||
Reference: https://blackforestlabs.ai/announcing-black-forest-labs/
|
||||
|
||||
Args:
|
||||
transformer ([`FluxTransformer2DModel`]):
|
||||
Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
|
||||
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`CLIPTextModel`]):
|
||||
[CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
|
||||
the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
|
||||
text_encoder_2 ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
|
||||
tokenizer (`CLIPTokenizer`):
|
||||
Tokenizer of class
|
||||
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
|
||||
tokenizer_2 (`T5TokenizerFast`):
|
||||
Second Tokenizer of class
|
||||
[T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->image_encoder->transformer->vae"
|
||||
_optional_components = ["image_encoder", "feature_extractor"]
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: CLIPTextModel,
|
||||
tokenizer: CLIPTokenizer,
|
||||
text_encoder_2: AutoModel,
|
||||
tokenizer_2: AutoTokenizer,
|
||||
transformer: FluxTransformer2DModel,
|
||||
image_encoder: CLIPVisionModelWithProjection = None,
|
||||
feature_extractor: CLIPImageProcessor = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer=tokenizer,
|
||||
tokenizer_2=tokenizer_2,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
image_encoder=image_encoder,
|
||||
feature_extractor=feature_extractor,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) if getattr(self, "vae", None) else 8
|
||||
# Flux latents are turned into 2x2 patches and packed. This means the latent width and height has to be divisible
|
||||
# by the patch size. So the vae scale factor is multiplied by the patch size to account for this
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * 2)
|
||||
self.tokenizer_max_length = (
|
||||
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
|
||||
)
|
||||
self.default_sample_size = 128
|
||||
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
|
||||
|
||||
# determine length of system prompt
|
||||
self.system_prompt_length = self.tokenizer_2(
|
||||
[self.system_prompt],
|
||||
padding="longest",
|
||||
return_tensors="pt",
|
||||
).input_ids[0].shape[0]
|
||||
|
||||
def _get_clip_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
num_images_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.tokenizer_max_length,
|
||||
truncation=True,
|
||||
return_overflowing_tokens=False,
|
||||
return_length=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, 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 = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {self.tokenizer_max_length} tokens: {removed_text}"
|
||||
)
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), output_hidden_states=False)
|
||||
|
||||
# Use pooled output of CLIPTextModel
|
||||
prompt_embeds = prompt_embeds.pooler_output
|
||||
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def _get_llm_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
max_sequence_length: int = 512,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt, self.tokenizer_2)
|
||||
|
||||
text_inputs = self.tokenizer_2(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length + self.system_prompt_length,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device)
|
||||
untruncated_ids = self.tokenizer_2(prompt, 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 = self.tokenizer_2.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {max_sequence_length + self.system_prompt_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
prompt_embeds = self.text_encoder_2(
|
||||
text_input_ids,
|
||||
attention_mask=prompt_attention_mask,
|
||||
output_hidden_states=True
|
||||
)
|
||||
prompt_embeds = prompt_embeds.hidden_states[-1]
|
||||
|
||||
# remove the system prompt from the input and attention mask
|
||||
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
|
||||
|
||||
dtype = self.text_encoder_2.dtype
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
|
||||
# duplicate text embeddings and attention mask 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(batch_size * num_images_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]],
|
||||
device: Optional[torch.device] = None,
|
||||
num_images_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
max_sequence_length: int = 512,
|
||||
lora_scale: Optional[float] = None,
|
||||
):
|
||||
r"""
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
used in all text-encoders
|
||||
device: (`torch.device`):
|
||||
torch device
|
||||
num_images_per_prompt (`int`):
|
||||
number of images that should be generated per prompt
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
lora_scale (`float`, *optional*):
|
||||
A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded.
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
# set lora scale so that monkey patched LoRA
|
||||
# function of text encoder can correctly access it
|
||||
if lora_scale is not None and isinstance(self, FluxLoraLoaderMixin):
|
||||
self._lora_scale = lora_scale
|
||||
|
||||
# dynamically adjust the LoRA scale
|
||||
if self.text_encoder is not None and USE_PEFT_BACKEND:
|
||||
scale_lora_layers(self.text_encoder, lora_scale)
|
||||
if self.text_encoder_2 is not None and USE_PEFT_BACKEND:
|
||||
scale_lora_layers(self.text_encoder_2, lora_scale)
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_2 = prompt_2 or prompt
|
||||
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
|
||||
|
||||
# We only use the pooled prompt output from the CLIPTextModel
|
||||
pooled_prompt_embeds = self._get_clip_prompt_embeds(
|
||||
prompt=prompt,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
)
|
||||
prompt_embeds = self._get_llm_prompt_embeds(
|
||||
prompt=prompt_2,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
if self.text_encoder is not None:
|
||||
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
|
||||
# Retrieve the original scale by scaling back the LoRA layers
|
||||
unscale_lora_layers(self.text_encoder, lora_scale)
|
||||
|
||||
if self.text_encoder_2 is not None:
|
||||
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
|
||||
# Retrieve the original scale by scaling back the LoRA layers
|
||||
unscale_lora_layers(self.text_encoder_2, lora_scale)
|
||||
|
||||
dtype = self.text_encoder.dtype if self.text_encoder is not None else self.transformer.dtype
|
||||
text_ids = torch.zeros(prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, text_ids
|
||||
|
||||
def encode_image(self, image, device, num_images_per_prompt):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.feature_extractor(image, return_tensors="pt").pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeds = self.image_encoder(image).image_embeds
|
||||
image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
return image_embeds
|
||||
|
||||
def prepare_ip_adapter_image_embeds(
|
||||
self, ip_adapter_image, ip_adapter_image_embeds, device, num_images_per_prompt
|
||||
):
|
||||
image_embeds = []
|
||||
if ip_adapter_image_embeds is None:
|
||||
if not isinstance(ip_adapter_image, list):
|
||||
ip_adapter_image = [ip_adapter_image]
|
||||
|
||||
if len(ip_adapter_image) != len(self.transformer.encoder_hid_proj.image_projection_layers):
|
||||
raise ValueError(
|
||||
f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(self.transformer.encoder_hid_proj.image_projection_layers)} IP Adapters."
|
||||
)
|
||||
|
||||
for single_ip_adapter_image, image_proj_layer in zip(
|
||||
ip_adapter_image, self.transformer.encoder_hid_proj.image_projection_layers
|
||||
):
|
||||
single_image_embeds = self.encode_image(single_ip_adapter_image, device, 1)
|
||||
|
||||
image_embeds.append(single_image_embeds[None, :])
|
||||
else:
|
||||
for single_image_embeds in ip_adapter_image_embeds:
|
||||
image_embeds.append(single_image_embeds)
|
||||
|
||||
ip_adapter_image_embeds = []
|
||||
for i, single_image_embeds in enumerate(image_embeds):
|
||||
single_image_embeds = torch.cat([single_image_embeds] * num_images_per_prompt, dim=0)
|
||||
single_image_embeds = single_image_embeds.to(device=device)
|
||||
ip_adapter_image_embeds.append(single_image_embeds)
|
||||
|
||||
return ip_adapter_image_embeds
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
negative_prompt=None,
|
||||
negative_prompt_2=None,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
pooled_prompt_embeds=None,
|
||||
negative_pooled_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
max_sequence_length=None,
|
||||
):
|
||||
if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0:
|
||||
logger.warning(
|
||||
f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly"
|
||||
)
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt_2 is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)):
|
||||
raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}")
|
||||
|
||||
if negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
elif negative_prompt_2 is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt_2`: {negative_prompt_2} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and pooled_prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`."
|
||||
)
|
||||
if negative_prompt_embeds is not None and negative_pooled_prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"If `negative_prompt_embeds` are provided, `negative_pooled_prompt_embeds` also have to be passed. Make sure to generate `negative_pooled_prompt_embeds` from the same text encoder that was used to generate `negative_prompt_embeds`."
|
||||
)
|
||||
|
||||
if max_sequence_length is not None and max_sequence_length > 512:
|
||||
raise ValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}")
|
||||
|
||||
@staticmethod
|
||||
def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
|
||||
latent_image_ids = torch.zeros(height, width, 3)
|
||||
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
|
||||
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
|
||||
|
||||
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
|
||||
|
||||
latent_image_ids = latent_image_ids.reshape(
|
||||
latent_image_id_height * latent_image_id_width, latent_image_id_channels
|
||||
)
|
||||
|
||||
return latent_image_ids.to(device=device, dtype=dtype)
|
||||
|
||||
@staticmethod
|
||||
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
|
||||
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
|
||||
latents = latents.permute(0, 2, 4, 1, 3, 5)
|
||||
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
|
||||
|
||||
return latents
|
||||
|
||||
@staticmethod
|
||||
def _unpack_latents(latents, height, width, vae_scale_factor):
|
||||
batch_size, num_patches, channels = latents.shape
|
||||
|
||||
# VAE applies 8x compression on images but we must also account for packing which requires
|
||||
# latent height and width to be divisible by 2.
|
||||
height = 2 * (int(height) // (vae_scale_factor * 2))
|
||||
width = 2 * (int(width) // (vae_scale_factor * 2))
|
||||
|
||||
latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2)
|
||||
latents = latents.permute(0, 3, 1, 4, 2, 5)
|
||||
|
||||
latents = latents.reshape(batch_size, channels // (2 * 2), height, width)
|
||||
|
||||
return latents
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
dtype,
|
||||
device,
|
||||
generator,
|
||||
latents=None,
|
||||
):
|
||||
# VAE applies 8x compression on images but we must also account for packing which requires
|
||||
# latent height and width to be divisible by 2.
|
||||
height = 2 * (int(height) // (self.vae_scale_factor * 2))
|
||||
width = 2 * (int(width) // (self.vae_scale_factor * 2))
|
||||
|
||||
shape = (batch_size, num_channels_latents, height, width)
|
||||
|
||||
if latents is not None:
|
||||
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
|
||||
return latents.to(device=device, dtype=dtype), latent_image_ids
|
||||
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width)
|
||||
|
||||
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
|
||||
|
||||
return latents, latent_image_ids
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def joint_attention_kwargs(self):
|
||||
return self._joint_attention_kwargs
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def current_timestep(self):
|
||||
return self._current_timestep
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
true_cfg_scale: float = 1.0,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
guidance_scale: float = 3.5,
|
||||
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,
|
||||
ip_adapter_image: Optional[PipelineImageInput] = None,
|
||||
ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
|
||||
negative_ip_adapter_image: Optional[PipelineImageInput] = None,
|
||||
negative_ip_adapter_image_embeds: Optional[List[torch.Tensor]] = 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,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead.
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `true_cfg_scale` is
|
||||
not greater than `1`).
|
||||
negative_prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
|
||||
`text_encoder_2`. If not defined, `negative_prompt` is used in all the text-encoders.
|
||||
true_cfg_scale (`float`, *optional*, defaults to 1.0):
|
||||
When > 1.0 and a provided `negative_prompt`, enables true classifier-free guidance.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, *optional*, defaults to 7.0):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
|
||||
ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
|
||||
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
|
||||
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
|
||||
provided, embeddings are computed from the `ip_adapter_image` input argument.
|
||||
negative_ip_adapter_image:
|
||||
(`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
|
||||
negative_ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
|
||||
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
|
||||
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
|
||||
provided, embeddings are computed from the `ip_adapter_image` input argument.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
|
||||
input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
|
||||
joint_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
|
||||
is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
|
||||
images.
|
||||
"""
|
||||
|
||||
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,
|
||||
negative_prompt=negative_prompt,
|
||||
negative_prompt_2=negative_prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
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._current_timestep = None
|
||||
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
|
||||
)
|
||||
has_neg_prompt = negative_prompt is not None or (
|
||||
negative_prompt_embeds is not None and negative_pooled_prompt_embeds is not None
|
||||
)
|
||||
do_true_cfg = true_cfg_scale > 1 and has_neg_prompt
|
||||
(
|
||||
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,
|
||||
)
|
||||
if do_true_cfg:
|
||||
(
|
||||
negative_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
_,
|
||||
) = 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) if sigmas is None else sigmas
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.get("base_image_seq_len", 256),
|
||||
self.scheduler.config.get("max_image_seq_len", 4096),
|
||||
self.scheduler.config.get("base_shift", 0.5),
|
||||
self.scheduler.config.get("max_shift", 1.16),
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# handle guidance
|
||||
if self.transformer.config.guidance_embeds:
|
||||
guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
if (ip_adapter_image is not None or ip_adapter_image_embeds is not None) and (
|
||||
negative_ip_adapter_image is None and negative_ip_adapter_image_embeds is None
|
||||
):
|
||||
negative_ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
|
||||
elif (ip_adapter_image is None and ip_adapter_image_embeds is None) and (
|
||||
negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None
|
||||
):
|
||||
ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
|
||||
|
||||
if self.joint_attention_kwargs is None:
|
||||
self._joint_attention_kwargs = {}
|
||||
|
||||
image_embeds = None
|
||||
negative_image_embeds = None
|
||||
if ip_adapter_image is not None or ip_adapter_image_embeds is not None:
|
||||
image_embeds = self.prepare_ip_adapter_image_embeds(
|
||||
ip_adapter_image,
|
||||
ip_adapter_image_embeds,
|
||||
device,
|
||||
batch_size * num_images_per_prompt,
|
||||
)
|
||||
if negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None:
|
||||
negative_image_embeds = self.prepare_ip_adapter_image_embeds(
|
||||
negative_ip_adapter_image,
|
||||
negative_ip_adapter_image_embeds,
|
||||
device,
|
||||
batch_size * num_images_per_prompt,
|
||||
)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
if image_embeds is not None:
|
||||
self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
noise_pred = 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]
|
||||
|
||||
if do_true_cfg:
|
||||
if negative_image_embeds is not None:
|
||||
self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds
|
||||
neg_noise_pred = self.transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep / 1000,
|
||||
guidance=guidance,
|
||||
pooled_projections=negative_pooled_prompt_embeds,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
txt_ids=text_ids,
|
||||
img_ids=latent_image_ids,
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred)
|
||||
|
||||
# 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()
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
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)
|
||||
176
toolkit/models/flux.py
Normal file
176
toolkit/models/flux.py
Normal file
@@ -0,0 +1,176 @@
|
||||
|
||||
# forward that bypasses the guidance embedding so it can be avoided during training.
|
||||
from functools import partial
|
||||
from typing import Optional
|
||||
import torch
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
|
||||
def guidance_embed_bypass_forward(self, timestep, guidance, pooled_projection):
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
timesteps_emb = self.timestep_embedder(
|
||||
timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D)
|
||||
pooled_projections = self.text_embedder(pooled_projection)
|
||||
conditioning = timesteps_emb + pooled_projections
|
||||
return conditioning
|
||||
|
||||
# bypass the forward function
|
||||
|
||||
|
||||
def bypass_flux_guidance(transformer):
|
||||
if hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
# dont bypass if it doesnt have the guidance embedding
|
||||
if not hasattr(transformer.time_text_embed, 'guidance_embedder'):
|
||||
return
|
||||
transformer.time_text_embed._bfg_orig_forward = transformer.time_text_embed.forward
|
||||
transformer.time_text_embed.forward = partial(
|
||||
guidance_embed_bypass_forward, transformer.time_text_embed
|
||||
)
|
||||
|
||||
# restore the forward function
|
||||
|
||||
|
||||
def restore_flux_guidance(transformer):
|
||||
if not hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
transformer.time_text_embed.forward = transformer.time_text_embed._bfg_orig_forward
|
||||
del transformer.time_text_embed._bfg_orig_forward
|
||||
|
||||
def new_device_to(self: FluxTransformer2DModel, *args, **kwargs):
|
||||
# Store original device if provided in args or kwargs
|
||||
device_in_kwargs = 'device' in kwargs
|
||||
device_in_args = any(isinstance(arg, (str, torch.device)) for arg in args)
|
||||
|
||||
device = None
|
||||
# Remove device from kwargs if present
|
||||
if device_in_kwargs:
|
||||
device = kwargs['device']
|
||||
del kwargs['device']
|
||||
|
||||
# Only filter args if we detected a device argument
|
||||
if device_in_args:
|
||||
args = list(args)
|
||||
for idx, arg in enumerate(args):
|
||||
if isinstance(arg, (str, torch.device)):
|
||||
device = arg
|
||||
del args[idx]
|
||||
|
||||
self.pos_embed = self.pos_embed.to(device, *args, **kwargs)
|
||||
self.time_text_embed = self.time_text_embed.to(device, *args, **kwargs)
|
||||
self.context_embedder = self.context_embedder.to(device, *args, **kwargs)
|
||||
self.x_embedder = self.x_embedder.to(device, *args, **kwargs)
|
||||
for block in self.transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
for block in self.single_transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
|
||||
self.norm_out = self.norm_out.to(device, *args, **kwargs)
|
||||
self.proj_out = self.proj_out.to(device, *args, **kwargs)
|
||||
|
||||
|
||||
|
||||
return self
|
||||
|
||||
|
||||
|
||||
|
||||
def split_gpu_double_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(self._split_device)
|
||||
if encoder_hidden_states.device != self._split_device:
|
||||
encoder_hidden_states = encoder_hidden_states.to(self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
return self._pre_gpu_split_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
|
||||
|
||||
def split_gpu_single_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
**kwargs
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(device=self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(device=self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
|
||||
hidden_state_out = self._pre_gpu_split_forward(hidden_states, temb, image_rotary_emb, joint_attention_kwargs, **kwargs)
|
||||
if hasattr(self, "_split_output_device"):
|
||||
return hidden_state_out.to(self._split_output_device)
|
||||
return hidden_state_out
|
||||
|
||||
|
||||
def add_model_gpu_splitter_to_flux(
|
||||
transformer: FluxTransformer2DModel,
|
||||
# ~ 5 billion for all other params
|
||||
other_module_params: Optional[int] = 5e9,
|
||||
# since they are not trainable, multiply by smaller number
|
||||
other_module_param_count_scale: Optional[float] = 0.3
|
||||
):
|
||||
gpu_id_list = [i for i in range(torch.cuda.device_count())]
|
||||
|
||||
# if len(gpu_id_list) > 2:
|
||||
# raise ValueError("Cannot split to more than 2 GPUs currently.")
|
||||
other_module_params *= other_module_param_count_scale
|
||||
|
||||
# since we are not tuning the
|
||||
total_params = sum(p.numel() for p in transformer.parameters()) + other_module_params
|
||||
|
||||
params_per_gpu = total_params / len(gpu_id_list)
|
||||
|
||||
current_gpu_idx = 0
|
||||
# text encoders, vae, and some non block layers will all be on gpu 0
|
||||
current_gpu_params = other_module_params
|
||||
|
||||
for double_block in transformer.transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
double_block._pre_gpu_split_forward = double_block.forward
|
||||
double_block.forward = partial(
|
||||
split_gpu_double_block_forward, double_block)
|
||||
double_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in double_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
for single_block in transformer.single_transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
single_block._pre_gpu_split_forward = single_block.forward
|
||||
single_block.forward = partial(
|
||||
split_gpu_single_block_forward, single_block)
|
||||
single_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in single_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
# add output device to last layer
|
||||
transformer.single_transformer_blocks[-1]._split_output_device = torch.device("cuda:0")
|
||||
|
||||
transformer._pre_gpu_split_to = transformer.to
|
||||
transformer.to = partial(new_device_to, transformer)
|
||||
94
toolkit/models/flux_sage_attn.py
Normal file
94
toolkit/models/flux_sage_attn.py
Normal file
@@ -0,0 +1,94 @@
|
||||
from typing import Optional
|
||||
from diffusers.models.attention_processor import Attention
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class FluxSageAttnProcessor2_0:
|
||||
"""Attention processor used typically in processing the SD3-like self-attention projections."""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("FluxAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
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:
|
||||
from sageattention import sageattn
|
||||
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `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)
|
||||
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `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:
|
||||
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 = sageattn(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)
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
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)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
@@ -136,6 +136,8 @@ class InstantLoRAMidModule(torch.nn.Module):
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
@@ -170,6 +172,8 @@ class InstantLoRAMidModule(torch.nn.Module):
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
@@ -211,7 +215,8 @@ class InstantLoRAModule(torch.nn.Module):
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion'
|
||||
sd: 'StableDiffusion',
|
||||
config=None
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
|
||||
191
toolkit/models/llm_adapter.py
Normal file
191
toolkit/models/llm_adapter.py
Normal file
@@ -0,0 +1,191 @@
|
||||
from functools import partial
|
||||
import sys
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union, TYPE_CHECKING
|
||||
|
||||
from diffusers.models.transformers.transformer_flux import FluxTransformerBlock
|
||||
from transformers import AutoModel, AutoTokenizer, Qwen2Model, LlamaModel, Qwen2Tokenizer, LlamaTokenizer
|
||||
|
||||
from toolkit import train_tools
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from diffusers import Transformer2DModel
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion, PixArtSigmaPipeline
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
LLM = Union[Qwen2Model, LlamaModel]
|
||||
LLMTokenizer = Union[Qwen2Tokenizer, LlamaTokenizer]
|
||||
|
||||
|
||||
def new_context_embedder_forward(self, x):
|
||||
if self._adapter_ref().is_active:
|
||||
x = self._context_embedder_ref()(x)
|
||||
else:
|
||||
x = self._orig_forward(x)
|
||||
return x
|
||||
|
||||
def new_block_forward(
|
||||
self: FluxTransformerBlock,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if self._adapter_ref().is_active:
|
||||
return self._new_block_ref()(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
else:
|
||||
return self._orig_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
|
||||
|
||||
class LLMAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
adapter: 'CustomAdapter',
|
||||
sd: 'StableDiffusion',
|
||||
llm: LLM,
|
||||
tokenizer: LLMTokenizer,
|
||||
num_cloned_blocks: int = 0,
|
||||
):
|
||||
super(LLMAdapter, self).__init__()
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.llm_ref: weakref.ref = weakref.ref(llm)
|
||||
self.tokenizer_ref: weakref.ref = weakref.ref(tokenizer)
|
||||
self.num_cloned_blocks = num_cloned_blocks
|
||||
self.apply_embedding_mask = False
|
||||
# make sure we can pad
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
# self.system_prompt = ""
|
||||
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
|
||||
|
||||
# determine length of system prompt
|
||||
sys_prompt_tokenized = tokenizer(
|
||||
[self.system_prompt],
|
||||
padding="longest",
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
sys_prompt_tokenized_ids = sys_prompt_tokenized.input_ids[0]
|
||||
|
||||
self.system_prompt_length = sys_prompt_tokenized_ids.shape[0]
|
||||
|
||||
print(f"System prompt length: {self.system_prompt_length}")
|
||||
|
||||
self.hidden_size = llm.config.hidden_size
|
||||
|
||||
blocks = []
|
||||
|
||||
if sd.is_flux:
|
||||
self.apply_embedding_mask = True
|
||||
self.context_embedder = nn.Linear(
|
||||
self.hidden_size, sd.unet.inner_dim)
|
||||
self.sequence_length = 512
|
||||
sd.unet.context_embedder._orig_forward = sd.unet.context_embedder.forward
|
||||
sd.unet.context_embedder.forward = partial(
|
||||
new_context_embedder_forward, sd.unet.context_embedder)
|
||||
sd.unet.context_embedder._context_embedder_ref = weakref.ref(self.context_embedder)
|
||||
# add a is active property to the context embedder
|
||||
sd.unet.context_embedder._adapter_ref = self.adapter_ref
|
||||
|
||||
for idx in range(self.num_cloned_blocks):
|
||||
block = FluxTransformerBlock(
|
||||
dim=sd.unet.inner_dim,
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
)
|
||||
# patch it in case it is quantized
|
||||
patch_dequantization_on_save(sd.unet.transformer_blocks[idx])
|
||||
state_dict = sd.unet.transformer_blocks[idx].state_dict()
|
||||
for key, value in state_dict.items():
|
||||
block.state_dict()[key].copy_(value)
|
||||
blocks.append(block)
|
||||
orig_block = sd.unet.transformer_blocks[idx]
|
||||
orig_block._orig_forward = orig_block.forward
|
||||
orig_block.forward = partial(
|
||||
new_block_forward, orig_block)
|
||||
orig_block._new_block_ref = weakref.ref(block)
|
||||
orig_block._adapter_ref = self.adapter_ref
|
||||
|
||||
elif sd.is_lumina2:
|
||||
self.context_embedder = nn.Linear(
|
||||
self.hidden_size, sd.unet.hidden_size)
|
||||
self.sequence_length = 256
|
||||
else:
|
||||
raise ValueError(
|
||||
"llm adapter currently only supports flux or lumina2")
|
||||
|
||||
self.blocks = nn.ModuleList(blocks)
|
||||
|
||||
def _get_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
max_sequence_length: int = 256,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
tokenizer = self.tokenizer_ref()
|
||||
text_encoder = self.llm_ref()
|
||||
device = text_encoder.device
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
text_inputs = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length + self.system_prompt_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids.to(device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device)
|
||||
|
||||
# remove the system prompt from the input and attention mask
|
||||
|
||||
prompt_embeds = text_encoder(
|
||||
text_input_ids, attention_mask=prompt_attention_mask, output_hidden_states=True
|
||||
)
|
||||
prompt_embeds = prompt_embeds.hidden_states[-1]
|
||||
|
||||
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
|
||||
|
||||
dtype = text_encoder.dtype
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
# make a getter to see if is active
|
||||
|
||||
@property
|
||||
def is_active(self):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def encode_text(self, prompt):
|
||||
|
||||
prompt = prompt if isinstance(prompt, list) else [prompt]
|
||||
|
||||
prompt = [self.system_prompt + p for p in prompt]
|
||||
# prompt = [self.system_prompt + p for p in prompt]
|
||||
|
||||
prompt_embeds, prompt_attention_mask = self._get_prompt_embeds(
|
||||
prompt=prompt,
|
||||
max_sequence_length=self.sequence_length,
|
||||
)
|
||||
|
||||
prompt_embeds = PromptEmbeds(
|
||||
prompt_embeds,
|
||||
attention_mask=prompt_attention_mask,
|
||||
).detach()
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def forward(self, input):
|
||||
return input
|
||||
331
toolkit/models/lokr.py
Normal file
331
toolkit/models/lokr.py
Normal file
@@ -0,0 +1,331 @@
|
||||
# based heavily on https://github.com/KohakuBlueleaf/LyCORIS/blob/eb460098187f752a5d66406d3affade6f0a07ece/lycoris/modules/lokr.py
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from toolkit.network_mixins import ToolkitModuleMixin
|
||||
|
||||
from typing import TYPE_CHECKING, Union, List
|
||||
|
||||
from optimum.quanto import QBytesTensor, QTensor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
from toolkit.lora_special import LoRASpecialNetwork
|
||||
|
||||
|
||||
def factorization(dimension: int, factor: int = -1) -> tuple[int, int]:
|
||||
'''
|
||||
return a tuple of two value of input dimension decomposed by the number closest to factor
|
||||
second value is higher or equal than first value.
|
||||
|
||||
In LoRA with Kroneckor Product, first value is a value for weight scale.
|
||||
secon value is a value for weight.
|
||||
|
||||
Becuase of non-commutative property, A⊗B ≠ B⊗A. Meaning of two matrices is slightly different.
|
||||
|
||||
examples)
|
||||
factor
|
||||
-1 2 4 8 16 ...
|
||||
127 -> 127, 1 127 -> 127, 1 127 -> 127, 1 127 -> 127, 1 127 -> 127, 1
|
||||
128 -> 16, 8 128 -> 64, 2 128 -> 32, 4 128 -> 16, 8 128 -> 16, 8
|
||||
250 -> 125, 2 250 -> 125, 2 250 -> 125, 2 250 -> 125, 2 250 -> 125, 2
|
||||
360 -> 45, 8 360 -> 180, 2 360 -> 90, 4 360 -> 45, 8 360 -> 45, 8
|
||||
512 -> 32, 16 512 -> 256, 2 512 -> 128, 4 512 -> 64, 8 512 -> 32, 16
|
||||
1024 -> 32, 32 1024 -> 512, 2 1024 -> 256, 4 1024 -> 128, 8 1024 -> 64, 16
|
||||
'''
|
||||
|
||||
if factor > 0 and (dimension % factor) == 0:
|
||||
m = factor
|
||||
n = dimension // factor
|
||||
return m, n
|
||||
if factor == -1:
|
||||
factor = dimension
|
||||
m, n = 1, dimension
|
||||
length = m + n
|
||||
while m < n:
|
||||
new_m = m + 1
|
||||
while dimension % new_m != 0:
|
||||
new_m += 1
|
||||
new_n = dimension // new_m
|
||||
if new_m + new_n > length or new_m > factor:
|
||||
break
|
||||
else:
|
||||
m, n = new_m, new_n
|
||||
if m > n:
|
||||
n, m = m, n
|
||||
return m, n
|
||||
|
||||
|
||||
def make_weight_cp(t, wa, wb):
|
||||
rebuild2 = torch.einsum('i j k l, i p, j r -> p r k l',
|
||||
t, wa, wb) # [c, d, k1, k2]
|
||||
return rebuild2
|
||||
|
||||
|
||||
def make_kron(w1, w2, scale):
|
||||
if len(w2.shape) == 4:
|
||||
w1 = w1.unsqueeze(2).unsqueeze(2)
|
||||
w2 = w2.contiguous()
|
||||
rebuild = torch.kron(w1, w2)
|
||||
|
||||
return rebuild*scale
|
||||
|
||||
|
||||
class LokrModule(ToolkitModuleMixin, nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
lora_name,
|
||||
org_module: nn.Module,
|
||||
multiplier=1.0,
|
||||
lora_dim=4,
|
||||
alpha=1,
|
||||
dropout=0.,
|
||||
rank_dropout=0.,
|
||||
module_dropout=0.,
|
||||
use_cp=False,
|
||||
decompose_both=False,
|
||||
network: 'LoRASpecialNetwork' = None,
|
||||
factor: int = -1, # factorization factor
|
||||
**kwargs,
|
||||
):
|
||||
""" if alpha == 0 or None, alpha is rank (no scaling). """
|
||||
ToolkitModuleMixin.__init__(self, network=network)
|
||||
torch.nn.Module.__init__(self)
|
||||
factor = int(factor)
|
||||
self.lora_name = lora_name
|
||||
self.lora_dim = lora_dim
|
||||
self.cp = False
|
||||
self.use_w1 = False
|
||||
self.use_w2 = False
|
||||
self.can_merge_in = True
|
||||
|
||||
self.shape = org_module.weight.shape
|
||||
if org_module.__class__.__name__ == 'Conv2d':
|
||||
in_dim = org_module.in_channels
|
||||
k_size = org_module.kernel_size
|
||||
out_dim = org_module.out_channels
|
||||
|
||||
in_m, in_n = factorization(in_dim, factor)
|
||||
out_l, out_k = factorization(out_dim, factor)
|
||||
# ((a, b), (c, d), *k_size)
|
||||
shape = ((out_l, out_k), (in_m, in_n), *k_size)
|
||||
|
||||
self.cp = use_cp and k_size != (1, 1)
|
||||
if decompose_both and lora_dim < max(shape[0][0], shape[1][0])/2:
|
||||
self.lokr_w1_a = nn.Parameter(
|
||||
torch.empty(shape[0][0], lora_dim))
|
||||
self.lokr_w1_b = nn.Parameter(
|
||||
torch.empty(lora_dim, shape[1][0]))
|
||||
else:
|
||||
self.use_w1 = True
|
||||
self.lokr_w1 = nn.Parameter(torch.empty(
|
||||
shape[0][0], shape[1][0])) # a*c, 1-mode
|
||||
|
||||
if lora_dim >= max(shape[0][1], shape[1][1])/2:
|
||||
self.use_w2 = True
|
||||
self.lokr_w2 = nn.Parameter(torch.empty(
|
||||
shape[0][1], shape[1][1], *k_size))
|
||||
elif self.cp:
|
||||
self.lokr_t2 = nn.Parameter(torch.empty(
|
||||
lora_dim, lora_dim, shape[2], shape[3]))
|
||||
self.lokr_w2_a = nn.Parameter(
|
||||
torch.empty(lora_dim, shape[0][1])) # b, 1-mode
|
||||
self.lokr_w2_b = nn.Parameter(
|
||||
torch.empty(lora_dim, shape[1][1])) # d, 2-mode
|
||||
else: # Conv2d not cp
|
||||
# bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2]
|
||||
self.lokr_w2_a = nn.Parameter(
|
||||
torch.empty(shape[0][1], lora_dim))
|
||||
self.lokr_w2_b = nn.Parameter(torch.empty(
|
||||
lora_dim, shape[1][1]*shape[2]*shape[3]))
|
||||
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2)
|
||||
|
||||
self.op = F.conv2d
|
||||
self.extra_args = {
|
||||
"stride": org_module.stride,
|
||||
"padding": org_module.padding,
|
||||
"dilation": org_module.dilation,
|
||||
"groups": org_module.groups
|
||||
}
|
||||
|
||||
else: # Linear
|
||||
in_dim = org_module.in_features
|
||||
out_dim = org_module.out_features
|
||||
|
||||
in_m, in_n = factorization(in_dim, factor)
|
||||
out_l, out_k = factorization(out_dim, factor)
|
||||
# ((a, b), (c, d)), out_dim = a*c, in_dim = b*d
|
||||
shape = ((out_l, out_k), (in_m, in_n))
|
||||
|
||||
# smaller part. weight scale
|
||||
if decompose_both and lora_dim < max(shape[0][0], shape[1][0])/2:
|
||||
self.lokr_w1_a = nn.Parameter(
|
||||
torch.empty(shape[0][0], lora_dim))
|
||||
self.lokr_w1_b = nn.Parameter(
|
||||
torch.empty(lora_dim, shape[1][0]))
|
||||
else:
|
||||
self.use_w1 = True
|
||||
self.lokr_w1 = nn.Parameter(torch.empty(
|
||||
shape[0][0], shape[1][0])) # a*c, 1-mode
|
||||
|
||||
if lora_dim < max(shape[0][1], shape[1][1])/2:
|
||||
# bigger part. weight and LoRA. [b, dim] x [dim, d]
|
||||
self.lokr_w2_a = nn.Parameter(
|
||||
torch.empty(shape[0][1], lora_dim))
|
||||
self.lokr_w2_b = nn.Parameter(
|
||||
torch.empty(lora_dim, shape[1][1]))
|
||||
# w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd)
|
||||
else:
|
||||
self.use_w2 = True
|
||||
self.lokr_w2 = nn.Parameter(
|
||||
torch.empty(shape[0][1], shape[1][1]))
|
||||
|
||||
self.op = F.linear
|
||||
self.extra_args = {}
|
||||
|
||||
self.dropout = dropout
|
||||
if dropout:
|
||||
print("[WARN]LoKr haven't implemented normal dropout yet.")
|
||||
self.rank_dropout = rank_dropout
|
||||
self.module_dropout = module_dropout
|
||||
|
||||
if isinstance(alpha, torch.Tensor):
|
||||
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
|
||||
alpha = lora_dim if alpha is None or alpha == 0 else alpha
|
||||
if self.use_w2 and self.use_w1:
|
||||
# use scale = 1
|
||||
alpha = lora_dim
|
||||
self.scale = alpha / self.lora_dim
|
||||
self.register_buffer('alpha', torch.tensor(alpha)) # treat as constant
|
||||
|
||||
if self.use_w2:
|
||||
torch.nn.init.constant_(self.lokr_w2, 0)
|
||||
else:
|
||||
if self.cp:
|
||||
torch.nn.init.kaiming_uniform_(self.lokr_t2, a=math.sqrt(5))
|
||||
torch.nn.init.kaiming_uniform_(self.lokr_w2_a, a=math.sqrt(5))
|
||||
torch.nn.init.constant_(self.lokr_w2_b, 0)
|
||||
|
||||
if self.use_w1:
|
||||
torch.nn.init.kaiming_uniform_(self.lokr_w1, a=math.sqrt(5))
|
||||
else:
|
||||
torch.nn.init.kaiming_uniform_(self.lokr_w1_a, a=math.sqrt(5))
|
||||
torch.nn.init.kaiming_uniform_(self.lokr_w1_b, a=math.sqrt(5))
|
||||
|
||||
self.multiplier = multiplier
|
||||
self.org_module = [org_module]
|
||||
weight = make_kron(
|
||||
self.lokr_w1 if self.use_w1 else self.lokr_w1_a@self.lokr_w1_b,
|
||||
(self.lokr_w2 if self.use_w2
|
||||
else make_weight_cp(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b) if self.cp
|
||||
else self.lokr_w2_a@self.lokr_w2_b),
|
||||
torch.tensor(self.multiplier * self.scale)
|
||||
)
|
||||
assert torch.sum(torch.isnan(weight)) == 0, "weight is nan"
|
||||
|
||||
# Same as locon.py
|
||||
def apply_to(self):
|
||||
self.org_forward = self.org_module[0].forward
|
||||
self.org_module[0].forward = self.forward
|
||||
|
||||
def get_weight(self, orig_weight=None):
|
||||
weight = make_kron(
|
||||
self.lokr_w1 if self.use_w1 else self.lokr_w1_a@self.lokr_w1_b,
|
||||
(self.lokr_w2 if self.use_w2
|
||||
else make_weight_cp(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b) if self.cp
|
||||
else self.lokr_w2_a@self.lokr_w2_b),
|
||||
torch.tensor(self.scale)
|
||||
)
|
||||
if orig_weight is not None:
|
||||
weight = weight.reshape(orig_weight.shape)
|
||||
if self.training and self.rank_dropout:
|
||||
drop = torch.rand(weight.size(0)) < self.rank_dropout
|
||||
weight *= drop.view(-1, [1] *
|
||||
len(weight.shape[1:])).to(weight.device)
|
||||
return weight
|
||||
|
||||
@torch.no_grad()
|
||||
def merge_in(self, merge_weight=1.0):
|
||||
if not self.can_merge_in:
|
||||
return
|
||||
|
||||
# extract weight from org_module
|
||||
org_sd = self.org_module[0].state_dict()
|
||||
# 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()
|
||||
|
||||
scale = self.scale
|
||||
# handle trainable scaler method locon does
|
||||
if hasattr(self, 'scalar'):
|
||||
scale = scale * self.scalar
|
||||
|
||||
lokr_weight = self.get_weight(weight)
|
||||
|
||||
merged_weight = (
|
||||
weight
|
||||
+ (lokr_weight * merge_weight).to(weight.device, dtype=weight.dtype)
|
||||
)
|
||||
|
||||
# set weight to org_module
|
||||
org_sd[weight_key] = merged_weight.to(orig_dtype)
|
||||
self.org_module[0].load_state_dict(org_sd)
|
||||
|
||||
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:
|
||||
if isinstance(self.org_module[0].bias, QTensor) or isinstance(self.org_module[0].bias, QBytesTensor):
|
||||
return self.org_module[0].bias.dequantize().data.detach()
|
||||
else:
|
||||
return self.org_module[0].bias.data.detach()
|
||||
return None
|
||||
|
||||
def _call_forward(self, x):
|
||||
if isinstance(x, QTensor) or isinstance(x, QBytesTensor):
|
||||
x = x.dequantize()
|
||||
|
||||
orig_dtype = x.dtype
|
||||
|
||||
orig_weight = self.get_orig_weight()
|
||||
lokr_weight = self.get_weight(orig_weight).to(dtype=orig_weight.dtype)
|
||||
multiplier = self.network_ref().torch_multiplier
|
||||
|
||||
if x.dtype != orig_weight.dtype:
|
||||
x = x.to(dtype=orig_weight.dtype)
|
||||
|
||||
# we do not currently support split batch multipliers for lokr. Just do a mean
|
||||
multiplier = torch.mean(multiplier)
|
||||
|
||||
weight = (
|
||||
orig_weight
|
||||
+ lokr_weight * multiplier
|
||||
)
|
||||
bias = self.get_orig_bias()
|
||||
if bias is not None:
|
||||
bias = bias.to(weight.device, dtype=weight.dtype)
|
||||
output = self.op(
|
||||
x,
|
||||
weight.view(self.shape),
|
||||
bias,
|
||||
**self.extra_args
|
||||
)
|
||||
return output.to(orig_dtype)
|
||||
567
toolkit/models/lumina2.py
Normal file
567
toolkit/models/lumina2.py
Normal file
@@ -0,0 +1,567 @@
|
||||
# Copyright 2024 Alpha-VLLM Authors and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
from diffusers.utils import logging
|
||||
from diffusers.models.attention import LuminaFeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps, apply_rotary_emb, get_1d_rotary_pos_embed
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import LuminaLayerNormContinuous, LuminaRMSNormZero, RMSNorm
|
||||
import torch
|
||||
from torch.profiler import profile, record_function, ProfilerActivity
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
do_profile = False
|
||||
|
||||
|
||||
class Lumina2CombinedTimestepCaptionEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 4096,
|
||||
cap_feat_dim: int = 2048,
|
||||
frequency_embedding_size: int = 256,
|
||||
norm_eps: float = 1e-5,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(
|
||||
num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0.0
|
||||
)
|
||||
|
||||
self.timestep_embedder = TimestepEmbedding(
|
||||
in_channels=frequency_embedding_size, time_embed_dim=min(hidden_size, 1024)
|
||||
)
|
||||
|
||||
self.caption_embedder = nn.Sequential(
|
||||
RMSNorm(cap_feat_dim, eps=norm_eps), nn.Linear(cap_feat_dim, hidden_size, bias=True)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
timestep_proj = self.time_proj(timestep).type_as(hidden_states)
|
||||
time_embed = self.timestep_embedder(timestep_proj)
|
||||
caption_embed = self.caption_embedder(encoder_hidden_states)
|
||||
return time_embed, caption_embed
|
||||
|
||||
|
||||
class Lumina2AttnProcessor2_0:
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). This is
|
||||
used in the Lumina2Transformer2DModel model. It applies normalization and RoPE on query and key vectors.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
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: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
base_sequence_length: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
|
||||
# Get Query-Key-Value Pair
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query_dim = query.shape[-1]
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = query_dim // attn.heads
|
||||
dtype = query.dtype
|
||||
|
||||
# Get key-value heads
|
||||
kv_heads = inner_dim // head_dim
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim)
|
||||
key = key.view(batch_size, -1, kv_heads, head_dim)
|
||||
value = value.view(batch_size, -1, kv_heads, head_dim)
|
||||
|
||||
# Apply Query-Key Norm if needed
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# Apply RoPE if needed
|
||||
if image_rotary_emb is not None:
|
||||
query = apply_rotary_emb(query, image_rotary_emb, use_real=False)
|
||||
key = apply_rotary_emb(key, image_rotary_emb, use_real=False)
|
||||
|
||||
query, key = query.to(dtype), key.to(dtype)
|
||||
|
||||
# Apply proportional attention if true
|
||||
if base_sequence_length is not None:
|
||||
softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale
|
||||
else:
|
||||
softmax_scale = attn.scale
|
||||
|
||||
# perform Grouped-qurey Attention (GQA)
|
||||
n_rep = attn.heads // kv_heads
|
||||
if n_rep >= 1:
|
||||
key = key.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3)
|
||||
value = value.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3)
|
||||
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.bool().view(batch_size, 1, 1, -1)
|
||||
attention_mask = attention_mask.expand(-1, attn.heads, sequence_length, -1)
|
||||
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, scale=softmax_scale
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Lumina2TransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
num_kv_heads: int,
|
||||
multiple_of: int,
|
||||
ffn_dim_multiplier: float,
|
||||
norm_eps: float,
|
||||
modulation: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.head_dim = dim // num_attention_heads
|
||||
self.modulation = modulation
|
||||
|
||||
self.attn = Attention(
|
||||
query_dim=dim,
|
||||
cross_attention_dim=None,
|
||||
dim_head=dim // num_attention_heads,
|
||||
qk_norm="rms_norm",
|
||||
heads=num_attention_heads,
|
||||
kv_heads=num_kv_heads,
|
||||
eps=1e-5,
|
||||
bias=False,
|
||||
out_bias=False,
|
||||
processor=Lumina2AttnProcessor2_0(),
|
||||
)
|
||||
|
||||
self.feed_forward = LuminaFeedForward(
|
||||
dim=dim,
|
||||
inner_dim=4 * dim,
|
||||
multiple_of=multiple_of,
|
||||
ffn_dim_multiplier=ffn_dim_multiplier,
|
||||
)
|
||||
|
||||
if modulation:
|
||||
self.norm1 = LuminaRMSNormZero(
|
||||
embedding_dim=dim,
|
||||
norm_eps=norm_eps,
|
||||
norm_elementwise_affine=True,
|
||||
)
|
||||
else:
|
||||
self.norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
self.norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm2 = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
temb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if self.modulation:
|
||||
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb)
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = hidden_states + gate_msa.unsqueeze(1).tanh() * self.norm2(attn_output)
|
||||
mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1)))
|
||||
hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output)
|
||||
else:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
hidden_states = hidden_states + self.norm2(attn_output)
|
||||
mlp_output = self.feed_forward(self.ffn_norm1(hidden_states))
|
||||
hidden_states = hidden_states + self.ffn_norm2(mlp_output)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Lumina2RotaryPosEmbed(nn.Module):
|
||||
def __init__(self, theta: int, axes_dim: List[int], axes_lens: List[int] = (300, 512, 512), patch_size: int = 2):
|
||||
super().__init__()
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
self.axes_lens = axes_lens
|
||||
self.patch_size = patch_size
|
||||
|
||||
self.freqs_cis = self._precompute_freqs_cis(axes_dim, axes_lens, theta)
|
||||
|
||||
def _precompute_freqs_cis(self, axes_dim: List[int], axes_lens: List[int], theta: int) -> List[torch.Tensor]:
|
||||
freqs_cis = []
|
||||
for i, (d, e) in enumerate(zip(axes_dim, axes_lens)):
|
||||
emb = get_1d_rotary_pos_embed(d, e, theta=self.theta, freqs_dtype=torch.float64)
|
||||
freqs_cis.append(emb)
|
||||
return freqs_cis
|
||||
|
||||
def _get_freqs_cis(self, ids: torch.Tensor) -> torch.Tensor:
|
||||
result = []
|
||||
for i in range(len(self.axes_dim)):
|
||||
freqs = self.freqs_cis[i].to(ids.device)
|
||||
index = ids[:, :, i : i + 1].repeat(1, 1, freqs.shape[-1]).to(torch.int64)
|
||||
result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index))
|
||||
return torch.cat(result, dim=-1)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor):
|
||||
batch_size = len(hidden_states)
|
||||
p_h = p_w = self.patch_size
|
||||
device = hidden_states[0].device
|
||||
|
||||
l_effective_cap_len = attention_mask.sum(dim=1).tolist()
|
||||
# TODO: this should probably be refactored because all subtensors of hidden_states will be of same shape
|
||||
img_sizes = [(img.size(1), img.size(2)) for img in hidden_states]
|
||||
l_effective_img_len = [(H // p_h) * (W // p_w) for (H, W) in img_sizes]
|
||||
|
||||
max_seq_len = max((cap_len + img_len for cap_len, img_len in zip(l_effective_cap_len, l_effective_img_len)))
|
||||
max_img_len = max(l_effective_img_len)
|
||||
|
||||
position_ids = torch.zeros(batch_size, max_seq_len, 3, dtype=torch.int32, device=device)
|
||||
|
||||
for i in range(batch_size):
|
||||
cap_len = l_effective_cap_len[i]
|
||||
img_len = l_effective_img_len[i]
|
||||
H, W = img_sizes[i]
|
||||
H_tokens, W_tokens = H // p_h, W // p_w
|
||||
assert H_tokens * W_tokens == img_len
|
||||
|
||||
position_ids[i, :cap_len, 0] = torch.arange(cap_len, dtype=torch.int32, device=device)
|
||||
position_ids[i, cap_len : cap_len + img_len, 0] = cap_len
|
||||
row_ids = (
|
||||
torch.arange(H_tokens, dtype=torch.int32, device=device).view(-1, 1).repeat(1, W_tokens).flatten()
|
||||
)
|
||||
col_ids = (
|
||||
torch.arange(W_tokens, dtype=torch.int32, device=device).view(1, -1).repeat(H_tokens, 1).flatten()
|
||||
)
|
||||
position_ids[i, cap_len : cap_len + img_len, 1] = row_ids
|
||||
position_ids[i, cap_len : cap_len + img_len, 2] = col_ids
|
||||
|
||||
freqs_cis = self._get_freqs_cis(position_ids)
|
||||
|
||||
cap_freqs_cis_shape = list(freqs_cis.shape)
|
||||
cap_freqs_cis_shape[1] = attention_mask.shape[1]
|
||||
cap_freqs_cis = torch.zeros(*cap_freqs_cis_shape, device=device, dtype=freqs_cis.dtype)
|
||||
|
||||
img_freqs_cis_shape = list(freqs_cis.shape)
|
||||
img_freqs_cis_shape[1] = max_img_len
|
||||
img_freqs_cis = torch.zeros(*img_freqs_cis_shape, device=device, dtype=freqs_cis.dtype)
|
||||
|
||||
for i in range(batch_size):
|
||||
cap_len = l_effective_cap_len[i]
|
||||
img_len = l_effective_img_len[i]
|
||||
cap_freqs_cis[i, :cap_len] = freqs_cis[i, :cap_len]
|
||||
img_freqs_cis[i, :img_len] = freqs_cis[i, cap_len : cap_len + img_len]
|
||||
|
||||
flat_hidden_states = []
|
||||
for i in range(batch_size):
|
||||
img = hidden_states[i]
|
||||
C, H, W = img.size()
|
||||
img = img.view(C, H // p_h, p_h, W // p_w, p_w).permute(1, 3, 2, 4, 0).flatten(2).flatten(0, 1)
|
||||
flat_hidden_states.append(img)
|
||||
hidden_states = flat_hidden_states
|
||||
padded_img_embed = torch.zeros(
|
||||
batch_size, max_img_len, hidden_states[0].shape[-1], device=device, dtype=hidden_states[0].dtype
|
||||
)
|
||||
padded_img_mask = torch.zeros(batch_size, max_img_len, dtype=torch.bool, device=device)
|
||||
for i in range(batch_size):
|
||||
padded_img_embed[i, : l_effective_img_len[i]] = hidden_states[i]
|
||||
padded_img_mask[i, : l_effective_img_len[i]] = True
|
||||
|
||||
return (
|
||||
padded_img_embed,
|
||||
padded_img_mask,
|
||||
img_sizes,
|
||||
l_effective_cap_len,
|
||||
l_effective_img_len,
|
||||
freqs_cis,
|
||||
cap_freqs_cis,
|
||||
img_freqs_cis,
|
||||
max_seq_len,
|
||||
)
|
||||
|
||||
|
||||
class Lumina2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
r"""
|
||||
Lumina2NextDiT: Diffusion model with a Transformer backbone.
|
||||
|
||||
Parameters:
|
||||
sample_size (`int`): The width of the latent images. This is fixed during training since
|
||||
it is used to learn a number of position embeddings.
|
||||
patch_size (`int`, *optional*, (`int`, *optional*, defaults to 2):
|
||||
The size of each patch in the image. This parameter defines the resolution of patches fed into the model.
|
||||
in_channels (`int`, *optional*, defaults to 4):
|
||||
The number of input channels for the model. Typically, this matches the number of channels in the input
|
||||
images.
|
||||
hidden_size (`int`, *optional*, defaults to 4096):
|
||||
The dimensionality of the hidden layers in the model. This parameter determines the width of the model's
|
||||
hidden representations.
|
||||
num_layers (`int`, *optional*, default to 32):
|
||||
The number of layers in the model. This defines the depth of the neural network.
|
||||
num_attention_heads (`int`, *optional*, defaults to 32):
|
||||
The number of attention heads in each attention layer. This parameter specifies how many separate attention
|
||||
mechanisms are used.
|
||||
num_kv_heads (`int`, *optional*, defaults to 8):
|
||||
The number of key-value heads in the attention mechanism, if different from the number of attention heads.
|
||||
If None, it defaults to num_attention_heads.
|
||||
multiple_of (`int`, *optional*, defaults to 256):
|
||||
A factor that the hidden size should be a multiple of. This can help optimize certain hardware
|
||||
configurations.
|
||||
ffn_dim_multiplier (`float`, *optional*):
|
||||
A multiplier for the dimensionality of the feed-forward network. If None, it uses a default value based on
|
||||
the model configuration.
|
||||
norm_eps (`float`, *optional*, defaults to 1e-5):
|
||||
A small value added to the denominator for numerical stability in normalization layers.
|
||||
scaling_factor (`float`, *optional*, defaults to 1.0):
|
||||
A scaling factor applied to certain parameters or layers in the model. This can be used for adjusting the
|
||||
overall scale of the model's operations.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_no_split_modules = ["Lumina2TransformerBlock"]
|
||||
_skip_layerwise_casting_patterns = ["x_embedder", "norm"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
sample_size: int = 128,
|
||||
patch_size: int = 2,
|
||||
in_channels: int = 16,
|
||||
out_channels: Optional[int] = None,
|
||||
hidden_size: int = 2304,
|
||||
num_layers: int = 26,
|
||||
num_refiner_layers: int = 2,
|
||||
num_attention_heads: int = 24,
|
||||
num_kv_heads: int = 8,
|
||||
multiple_of: int = 256,
|
||||
ffn_dim_multiplier: Optional[float] = None,
|
||||
norm_eps: float = 1e-5,
|
||||
scaling_factor: float = 1.0,
|
||||
axes_dim_rope: Tuple[int, int, int] = (32, 32, 32),
|
||||
axes_lens: Tuple[int, int, int] = (300, 512, 512),
|
||||
cap_feat_dim: int = 1024,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.out_channels = out_channels or in_channels
|
||||
|
||||
# 1. Positional, patch & conditional embeddings
|
||||
self.rope_embedder = Lumina2RotaryPosEmbed(
|
||||
theta=10000, axes_dim=axes_dim_rope, axes_lens=axes_lens, patch_size=patch_size
|
||||
)
|
||||
|
||||
self.x_embedder = nn.Linear(in_features=patch_size * patch_size * in_channels, out_features=hidden_size)
|
||||
|
||||
self.time_caption_embed = Lumina2CombinedTimestepCaptionEmbedding(
|
||||
hidden_size=hidden_size, cap_feat_dim=cap_feat_dim, norm_eps=norm_eps
|
||||
)
|
||||
|
||||
# 2. Noise and context refinement blocks
|
||||
self.noise_refiner = nn.ModuleList(
|
||||
[
|
||||
Lumina2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=True,
|
||||
)
|
||||
for _ in range(num_refiner_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.context_refiner = nn.ModuleList(
|
||||
[
|
||||
Lumina2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=False,
|
||||
)
|
||||
for _ in range(num_refiner_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
Lumina2TransformerBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
num_kv_heads,
|
||||
multiple_of,
|
||||
ffn_dim_multiplier,
|
||||
norm_eps,
|
||||
modulation=True,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = LuminaLayerNormContinuous(
|
||||
embedding_dim=hidden_size,
|
||||
conditioning_embedding_dim=min(hidden_size, 1024),
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
bias=True,
|
||||
out_dim=patch_size * patch_size * self.out_channels,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> Union[torch.Tensor, Transformer2DModelOutput]:
|
||||
|
||||
hidden_size = self.config.get("hidden_size", 2304)
|
||||
# pad or slice text encoder
|
||||
if encoder_hidden_states.shape[2] > hidden_size:
|
||||
encoder_hidden_states = encoder_hidden_states[:, :, :hidden_size]
|
||||
elif encoder_hidden_states.shape[2] < hidden_size:
|
||||
encoder_hidden_states = F.pad(encoder_hidden_states, (0, hidden_size - encoder_hidden_states.shape[2]))
|
||||
|
||||
batch_size = hidden_states.size(0)
|
||||
|
||||
if do_profile:
|
||||
prof = torch.profiler.profile(
|
||||
activities=[
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
],
|
||||
)
|
||||
|
||||
prof.start()
|
||||
|
||||
# 1. Condition, positional & patch embedding
|
||||
temb, encoder_hidden_states = self.time_caption_embed(hidden_states, timestep, encoder_hidden_states)
|
||||
|
||||
(
|
||||
hidden_states,
|
||||
hidden_mask,
|
||||
hidden_sizes,
|
||||
encoder_hidden_len,
|
||||
hidden_len,
|
||||
joint_rotary_emb,
|
||||
encoder_rotary_emb,
|
||||
hidden_rotary_emb,
|
||||
max_seq_len,
|
||||
) = self.rope_embedder(hidden_states, attention_mask)
|
||||
|
||||
hidden_states = self.x_embedder(hidden_states)
|
||||
|
||||
# 2. Context & noise refinement
|
||||
for layer in self.context_refiner:
|
||||
encoder_hidden_states = layer(encoder_hidden_states, attention_mask, encoder_rotary_emb)
|
||||
|
||||
for layer in self.noise_refiner:
|
||||
hidden_states = layer(hidden_states, hidden_mask, hidden_rotary_emb, temb)
|
||||
|
||||
# 3. Attention mask preparation
|
||||
mask = hidden_states.new_zeros(batch_size, max_seq_len, dtype=torch.bool)
|
||||
padded_hidden_states = hidden_states.new_zeros(batch_size, max_seq_len, self.config.hidden_size)
|
||||
for i in range(batch_size):
|
||||
cap_len = encoder_hidden_len[i]
|
||||
img_len = hidden_len[i]
|
||||
mask[i, : cap_len + img_len] = True
|
||||
padded_hidden_states[i, :cap_len] = encoder_hidden_states[i, :cap_len]
|
||||
padded_hidden_states[i, cap_len : cap_len + img_len] = hidden_states[i, :img_len]
|
||||
hidden_states = padded_hidden_states
|
||||
|
||||
# 4. Transformer blocks
|
||||
for layer in self.layers:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
hidden_states = self._gradient_checkpointing_func(layer, hidden_states, mask, joint_rotary_emb, temb)
|
||||
else:
|
||||
hidden_states = layer(hidden_states, mask, joint_rotary_emb, temb)
|
||||
|
||||
# 5. Output norm & projection & unpatchify
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
|
||||
height_tokens = width_tokens = self.config.patch_size
|
||||
output = []
|
||||
for i in range(len(hidden_sizes)):
|
||||
height, width = hidden_sizes[i]
|
||||
begin = encoder_hidden_len[i]
|
||||
end = begin + (height // height_tokens) * (width // width_tokens)
|
||||
output.append(
|
||||
hidden_states[i][begin:end]
|
||||
.view(height // height_tokens, width // width_tokens, height_tokens, width_tokens, self.out_channels)
|
||||
.permute(4, 0, 2, 1, 3)
|
||||
.flatten(3, 4)
|
||||
.flatten(1, 2)
|
||||
)
|
||||
output = torch.stack(output, dim=0)
|
||||
|
||||
if do_profile:
|
||||
torch.cuda.synchronize() # Make sure all CUDA ops are done
|
||||
prof.stop()
|
||||
|
||||
print("\n==== Profile Results ====")
|
||||
print(prof.key_averages().table(sort_by="cpu_time_total", row_limit=1000))
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
return Transformer2DModelOutput(sample=output)
|
||||
618
toolkit/models/pixtral_vision.py
Normal file
618
toolkit/models/pixtral_vision.py
Normal file
@@ -0,0 +1,618 @@
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Any, Union, TYPE_CHECKING
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from dataclasses import dataclass
|
||||
from huggingface_hub import snapshot_download
|
||||
from safetensors.torch import load_file
|
||||
import json
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from xformers.ops.fmha.attn_bias import BlockDiagonalMask
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def _norm(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim: int, hidden_dim: int, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
||||
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# type: ignore
|
||||
return self.w2(nn.functional.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
def repeat_kv(keys: torch.Tensor, values: torch.Tensor, repeats: int, dim: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
keys = torch.repeat_interleave(keys, repeats=repeats, dim=dim)
|
||||
values = torch.repeat_interleave(values, repeats=repeats, dim=dim)
|
||||
return keys, values
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
xq: torch.Tensor,
|
||||
xk: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
|
||||
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
|
||||
freqs_cis = freqs_cis[:, None, :]
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(-2)
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(-2)
|
||||
return xq_out.type_as(xq), xk_out.type_as(xk)
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
head_dim: int,
|
||||
n_kv_heads: int,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.n_heads: int = n_heads
|
||||
self.head_dim: int = head_dim
|
||||
self.n_kv_heads: int = n_kv_heads
|
||||
|
||||
self.repeats = self.n_heads // self.n_kv_heads
|
||||
|
||||
self.scale = self.head_dim ** -0.5
|
||||
|
||||
self.wq = nn.Linear(dim, n_heads * head_dim, bias=False)
|
||||
self.wk = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
|
||||
self.wv = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
|
||||
self.wo = nn.Linear(n_heads * head_dim, dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
cache: Optional[Any] = None,
|
||||
mask: Optional['BlockDiagonalMask'] = None,
|
||||
) -> torch.Tensor:
|
||||
from xformers.ops.fmha import memory_efficient_attention
|
||||
assert mask is None or cache is None
|
||||
seqlen_sum, _ = x.shape
|
||||
|
||||
xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
|
||||
xq = xq.view(seqlen_sum, self.n_heads, self.head_dim)
|
||||
xk = xk.view(seqlen_sum, self.n_kv_heads, self.head_dim)
|
||||
xv = xv.view(seqlen_sum, self.n_kv_heads, self.head_dim)
|
||||
xq, xk = apply_rotary_emb(xq, xk, freqs_cis=freqs_cis)
|
||||
|
||||
if cache is None:
|
||||
key, val = xk, xv
|
||||
elif cache.prefill:
|
||||
key, val = cache.interleave_kv(xk, xv)
|
||||
cache.update(xk, xv)
|
||||
else:
|
||||
cache.update(xk, xv)
|
||||
key, val = cache.key, cache.value
|
||||
key = key.view(seqlen_sum * cache.max_seq_len,
|
||||
self.n_kv_heads, self.head_dim)
|
||||
val = val.view(seqlen_sum * cache.max_seq_len,
|
||||
self.n_kv_heads, self.head_dim)
|
||||
|
||||
# Repeat keys and values to match number of query heads
|
||||
key, val = repeat_kv(key, val, self.repeats, dim=1)
|
||||
|
||||
# xformers requires (B=1, S, H, D)
|
||||
xq, key, val = xq[None, ...], key[None, ...], val[None, ...]
|
||||
output = memory_efficient_attention(
|
||||
xq, key, val, mask if cache is None else cache.mask)
|
||||
output = output.view(seqlen_sum, self.n_heads * self.head_dim)
|
||||
|
||||
assert isinstance(output, torch.Tensor)
|
||||
|
||||
return self.wo(output) # type: ignore
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
norm_eps: float,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_heads = n_heads
|
||||
self.dim = dim
|
||||
self.attention = Attention(
|
||||
dim=dim,
|
||||
n_heads=n_heads,
|
||||
head_dim=head_dim,
|
||||
n_kv_heads=n_kv_heads,
|
||||
)
|
||||
self.attention_norm = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
self.feed_forward: nn.Module
|
||||
self.feed_forward = FeedForward(dim=dim, hidden_dim=hidden_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
cache: Optional[Any] = None,
|
||||
mask: Optional['BlockDiagonalMask'] = None,
|
||||
) -> torch.Tensor:
|
||||
r = self.attention.forward(self.attention_norm(x), freqs_cis, cache)
|
||||
h = x + r
|
||||
r = self.feed_forward.forward(self.ffn_norm(h))
|
||||
out = h + r
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class VisionEncoderArgs:
|
||||
hidden_size: int
|
||||
num_channels: int
|
||||
image_size: int
|
||||
patch_size: int
|
||||
intermediate_size: int
|
||||
num_hidden_layers: int
|
||||
num_attention_heads: int
|
||||
rope_theta: float = 1e4 # for rope-2D
|
||||
image_token_id: int = 10
|
||||
|
||||
|
||||
def precompute_freqs_cis_2d(
|
||||
dim: int,
|
||||
height: int,
|
||||
width: int,
|
||||
theta: float,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
freqs_cis: 2D complex tensor of shape (height, width, dim // 2) to be indexed by
|
||||
(height, width) position tuples
|
||||
"""
|
||||
# (dim / 2) frequency bases
|
||||
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
|
||||
|
||||
h = torch.arange(height, device=freqs.device)
|
||||
w = torch.arange(width, device=freqs.device)
|
||||
|
||||
freqs_h = torch.outer(h, freqs[::2]).float()
|
||||
freqs_w = torch.outer(w, freqs[1::2]).float()
|
||||
freqs_2d = torch.cat(
|
||||
[
|
||||
freqs_h[:, None, :].repeat(1, width, 1),
|
||||
freqs_w[None, :, :].repeat(height, 1, 1),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
return torch.polar(torch.ones_like(freqs_2d), freqs_2d)
|
||||
|
||||
|
||||
def position_meshgrid(
|
||||
patch_embeds_list: list[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
positions = torch.cat(
|
||||
[
|
||||
torch.stack(
|
||||
torch.meshgrid(
|
||||
torch.arange(p.shape[-2]),
|
||||
torch.arange(p.shape[-1]),
|
||||
indexing="ij",
|
||||
),
|
||||
dim=-1,
|
||||
).reshape(-1, 2)
|
||||
for p in patch_embeds_list
|
||||
]
|
||||
)
|
||||
return positions
|
||||
|
||||
|
||||
class PixtralVisionEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 1024,
|
||||
num_channels: int = 3,
|
||||
image_size: int = 1024,
|
||||
patch_size: int = 16,
|
||||
intermediate_size: int = 4096,
|
||||
num_hidden_layers: int = 24,
|
||||
num_attention_heads: int = 16,
|
||||
rope_theta: float = 1e4, # for rope-2D
|
||||
image_token_id: int = 10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.args = VisionEncoderArgs(
|
||||
hidden_size=hidden_size,
|
||||
num_channels=num_channels,
|
||||
image_size=image_size,
|
||||
patch_size=patch_size,
|
||||
intermediate_size=intermediate_size,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
num_attention_heads=num_attention_heads,
|
||||
rope_theta=rope_theta,
|
||||
image_token_id=image_token_id,
|
||||
)
|
||||
args = self.args
|
||||
self.patch_conv = nn.Conv2d(
|
||||
in_channels=args.num_channels,
|
||||
out_channels=args.hidden_size,
|
||||
kernel_size=args.patch_size,
|
||||
stride=args.patch_size,
|
||||
bias=False,
|
||||
)
|
||||
self.ln_pre = RMSNorm(args.hidden_size, eps=1e-5)
|
||||
self.transformer = VisionTransformerBlocks(args)
|
||||
|
||||
head_dim = self.args.hidden_size // self.args.num_attention_heads
|
||||
assert head_dim % 2 == 0, "ROPE requires even head_dim"
|
||||
self._freqs_cis: Optional[torch.Tensor] = None
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path: str) -> 'PixtralVisionEncoder':
|
||||
if os.path.isdir(pretrained_model_name_or_path):
|
||||
model_folder = pretrained_model_name_or_path
|
||||
else:
|
||||
model_folder = snapshot_download(pretrained_model_name_or_path)
|
||||
|
||||
# make sure there is a config
|
||||
if not os.path.exists(os.path.join(model_folder, "config.json")):
|
||||
raise ValueError(f"Could not find config.json in {model_folder}")
|
||||
|
||||
# load config
|
||||
with open(os.path.join(model_folder, "config.json"), "r") as f:
|
||||
config = json.load(f)
|
||||
|
||||
model = cls(**config)
|
||||
|
||||
# see if there is a state_dict
|
||||
if os.path.exists(os.path.join(model_folder, "model.safetensors")):
|
||||
state_dict = load_file(os.path.join(
|
||||
model_folder, "model.safetensors"))
|
||||
model.load_state_dict(state_dict)
|
||||
|
||||
return model
|
||||
|
||||
@property
|
||||
def max_patches_per_side(self) -> int:
|
||||
return self.args.image_size // self.args.patch_size
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def freqs_cis(self) -> torch.Tensor:
|
||||
if self._freqs_cis is None:
|
||||
self._freqs_cis = precompute_freqs_cis_2d(
|
||||
dim=self.args.hidden_size // self.args.num_attention_heads,
|
||||
height=self.max_patches_per_side,
|
||||
width=self.max_patches_per_side,
|
||||
theta=self.args.rope_theta,
|
||||
)
|
||||
|
||||
if self._freqs_cis.device != self.device:
|
||||
self._freqs_cis = self._freqs_cis.to(device=self.device)
|
||||
|
||||
return self._freqs_cis
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images: List[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
from xformers.ops.fmha.attn_bias import BlockDiagonalMask
|
||||
"""
|
||||
Args:
|
||||
images: list of N_img images of variable sizes, each of shape (C, H, W)
|
||||
|
||||
Returns:
|
||||
image_features: tensor of token features for all tokens of all images of
|
||||
shape (N_toks, D)
|
||||
"""
|
||||
assert isinstance(
|
||||
images, list), f"Expected list of images, got {type(images)}"
|
||||
assert all(len(img.shape) == 3 for img in
|
||||
images), f"Expected images with shape (C, H, W), got {[img.shape for img in images]}"
|
||||
# pass images through initial convolution independently
|
||||
patch_embeds_list = [self.patch_conv(
|
||||
img.unsqueeze(0)).squeeze(0) for img in images]
|
||||
|
||||
# flatten to a single sequence
|
||||
patch_embeds = torch.cat([p.flatten(1).permute(1, 0)
|
||||
for p in patch_embeds_list], dim=0)
|
||||
patch_embeds = self.ln_pre(patch_embeds)
|
||||
|
||||
# positional embeddings
|
||||
positions = position_meshgrid(patch_embeds_list).to(self.device)
|
||||
freqs_cis = self.freqs_cis[positions[:, 0], positions[:, 1]]
|
||||
|
||||
# pass through Transformer with a block diagonal mask delimiting images
|
||||
mask = BlockDiagonalMask.from_seqlens(
|
||||
[p.shape[-2] * p.shape[-1] for p in patch_embeds_list],
|
||||
)
|
||||
out = self.transformer(patch_embeds, mask=mask, freqs_cis=freqs_cis)
|
||||
|
||||
# remove batch dimension of the single sequence
|
||||
return out # type: ignore[no-any-return]
|
||||
|
||||
|
||||
class VisionLanguageAdapter(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int):
|
||||
super().__init__()
|
||||
self.w_in = nn.Linear(
|
||||
in_dim,
|
||||
out_dim,
|
||||
bias=True,
|
||||
)
|
||||
self.gelu = nn.GELU()
|
||||
self.w_out = nn.Linear(out_dim, out_dim, bias=True)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# type: ignore[no-any-return]
|
||||
return self.w_out(self.gelu(self.w_in(x)))
|
||||
|
||||
|
||||
class VisionTransformerBlocks(nn.Module):
|
||||
def __init__(self, args: VisionEncoderArgs):
|
||||
super().__init__()
|
||||
self.layers = torch.nn.ModuleList()
|
||||
for _ in range(args.num_hidden_layers):
|
||||
self.layers.append(
|
||||
TransformerBlock(
|
||||
dim=args.hidden_size,
|
||||
hidden_dim=args.intermediate_size,
|
||||
n_heads=args.num_attention_heads,
|
||||
n_kv_heads=args.num_attention_heads,
|
||||
head_dim=args.hidden_size // args.num_attention_heads,
|
||||
norm_eps=1e-5,
|
||||
)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
mask: 'BlockDiagonalMask',
|
||||
freqs_cis: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
for layer in self.layers:
|
||||
x = layer(x, mask=mask, freqs_cis=freqs_cis)
|
||||
return x
|
||||
|
||||
|
||||
DATASET_MEAN = [0.48145466, 0.4578275, 0.40821073] # RGB
|
||||
DATASET_STD = [0.26862954, 0.26130258, 0.27577711] # RGB
|
||||
|
||||
|
||||
def normalize(image: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Normalize a tensor image with mean and standard deviation.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): Image to be normalized, shape (C, H, W), values in [0, 1].
|
||||
mean (torch.Tensor): Mean for each channel.
|
||||
std (torch.Tensor): Standard deviation for each channel.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Normalized image with shape (C, H, W).
|
||||
"""
|
||||
assert image.shape[0] == len(mean) == len(
|
||||
std), f"{image.shape=}, {mean.shape=}, {std.shape=}"
|
||||
|
||||
# Reshape mean and std to (C, 1, 1) for broadcasting
|
||||
mean = mean.view(-1, 1, 1)
|
||||
std = std.view(-1, 1, 1)
|
||||
|
||||
return (image - mean) / std
|
||||
|
||||
|
||||
def transform_image(image: torch.Tensor, new_size: tuple[int, int]) -> torch.Tensor:
|
||||
"""
|
||||
Resize and normalize the input image.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): Input image tensor of shape (C, H, W), values in [0, 1].
|
||||
new_size (tuple[int, int]): Target size (height, width) for resizing.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Resized and normalized image tensor of shape (C, new_H, new_W).
|
||||
"""
|
||||
# Resize the image
|
||||
resized_image = torch.nn.functional.interpolate(
|
||||
image.unsqueeze(0),
|
||||
size=new_size,
|
||||
mode='bicubic',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
|
||||
# Normalize the image
|
||||
normalized_image = normalize(
|
||||
resized_image,
|
||||
torch.tensor(DATASET_MEAN, device=image.device, dtype=image.dtype),
|
||||
torch.tensor(DATASET_STD, device=image.device, dtype=image.dtype)
|
||||
)
|
||||
|
||||
return normalized_image
|
||||
|
||||
|
||||
class PixtralVisionImagePreprocessor:
|
||||
def __init__(self, image_patch_size=16, max_image_size=1024) -> None:
|
||||
self.image_patch_size = image_patch_size
|
||||
self.max_image_size = max_image_size
|
||||
self.image_token = 10
|
||||
|
||||
def _image_to_num_tokens(self, img: torch.Tensor, max_image_size = None) -> Tuple[int, int]:
|
||||
w: Union[int, float]
|
||||
h: Union[int, float]
|
||||
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
|
||||
w, h = img.shape[-1], img.shape[-2]
|
||||
|
||||
# originally, pixtral used the largest of the 2 dimensions, but we
|
||||
# will use the base size of the image based on number of pixels.
|
||||
# ratio = max(h / self.max_image_size, w / self.max_image_size) # original
|
||||
|
||||
base_size = int(math.sqrt(w * h))
|
||||
ratio = base_size / max_image_size
|
||||
if ratio > 1:
|
||||
w = round(w / ratio)
|
||||
h = round(h / ratio)
|
||||
|
||||
width_tokens = (w - 1) // self.image_patch_size + 1
|
||||
height_tokens = (h - 1) // self.image_patch_size + 1
|
||||
|
||||
return width_tokens, height_tokens
|
||||
|
||||
def __call__(self, image: torch.Tensor, max_image_size=None) -> torch.Tensor:
|
||||
"""
|
||||
Converts ImageChunks to numpy image arrays and image token ids
|
||||
|
||||
Args:
|
||||
image torch tensor with values 0-1 and shape of (C, H, W)
|
||||
|
||||
Returns:
|
||||
processed_image: tensor of token features for all tokens of all images of
|
||||
"""
|
||||
# should not have batch
|
||||
if len(image.shape) == 4:
|
||||
raise ValueError(
|
||||
f"Expected image with shape (C, H, W), got {image.shape}")
|
||||
|
||||
if image.min() < 0.0 or image.max() > 1.0:
|
||||
raise ValueError(
|
||||
f"image tensor values must be between 0 and 1. Got min: {image.min()}, max: {image.max()}")
|
||||
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
|
||||
w, h = self._image_to_num_tokens(image, max_image_size=max_image_size)
|
||||
assert w > 0
|
||||
assert h > 0
|
||||
|
||||
new_image_size = (
|
||||
w * self.image_patch_size,
|
||||
h * self.image_patch_size,
|
||||
)
|
||||
|
||||
processed_image = transform_image(image, new_image_size)
|
||||
|
||||
return processed_image
|
||||
|
||||
|
||||
class PixtralVisionImagePreprocessorCompatibleReturn:
|
||||
def __init__(self, pixel_values) -> None:
|
||||
self.pixel_values = pixel_values
|
||||
|
||||
|
||||
# Compatable version with ai toolkit flow
|
||||
class PixtralVisionImagePreprocessorCompatible(PixtralVisionImagePreprocessor):
|
||||
def __init__(self, image_patch_size=16, max_image_size=1024) -> None:
|
||||
super().__init__(
|
||||
image_patch_size=image_patch_size,
|
||||
max_image_size=max_image_size
|
||||
)
|
||||
self.size = {
|
||||
'height': max_image_size,
|
||||
'width': max_image_size
|
||||
}
|
||||
self.max_image_size = max_image_size
|
||||
self.image_mean = DATASET_MEAN
|
||||
self.image_std = DATASET_STD
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
images,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
max_image_size=None,
|
||||
) -> torch.Tensor:
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
out_stack = []
|
||||
if len(images.shape) == 3:
|
||||
images = images.unsqueeze(0)
|
||||
for i in range(images.shape[0]):
|
||||
image = images[i]
|
||||
processed_image = super().__call__(image, max_image_size=max_image_size)
|
||||
out_stack.append(processed_image)
|
||||
|
||||
output = torch.stack(out_stack, dim=0)
|
||||
return PixtralVisionImagePreprocessorCompatibleReturn(output)
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatibleReturn:
|
||||
def __init__(self, hidden_states) -> None:
|
||||
self.hidden_states = hidden_states
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatibleConfig:
|
||||
def __init__(self):
|
||||
self.image_size = 1024
|
||||
self.hidden_size = 1024
|
||||
self.patch_size = 16
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatible(PixtralVisionEncoder):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 1024,
|
||||
num_channels: int = 3,
|
||||
image_size: int = 1024,
|
||||
patch_size: int = 16,
|
||||
intermediate_size: int = 4096,
|
||||
num_hidden_layers: int = 24,
|
||||
num_attention_heads: int = 16,
|
||||
rope_theta: float = 1e4, # for rope-2D
|
||||
image_token_id: int = 10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(
|
||||
hidden_size=hidden_size,
|
||||
num_channels=num_channels,
|
||||
image_size=image_size,
|
||||
patch_size=patch_size,
|
||||
intermediate_size=intermediate_size,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
num_attention_heads=num_attention_heads,
|
||||
rope_theta=rope_theta,
|
||||
image_token_id=image_token_id,
|
||||
)
|
||||
self.config = PixtralVisionEncoderCompatibleConfig()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images,
|
||||
output_hidden_states=True,
|
||||
) -> torch.Tensor:
|
||||
out_stack = []
|
||||
if len(images.shape) == 3:
|
||||
images = images.unsqueeze(0)
|
||||
for i in range(images.shape[0]):
|
||||
image = images[i]
|
||||
# must be in an array
|
||||
image_output = super().forward([image])
|
||||
out_stack.append(image_output)
|
||||
|
||||
output = torch.stack(out_stack, dim=0)
|
||||
return PixtralVisionEncoderCompatibleReturn([output])
|
||||
26
toolkit/models/redux.py
Normal file
26
toolkit/models/redux.py
Normal file
@@ -0,0 +1,26 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ReduxImageEncoder(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
redux_dim: int = 1152,
|
||||
txt_in_features: int = 4096,
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.redux_dim = redux_dim
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.redux_up = nn.Linear(redux_dim, txt_in_features * 3, dtype=dtype)
|
||||
self.redux_down = nn.Linear(
|
||||
txt_in_features * 3, txt_in_features, dtype=dtype)
|
||||
|
||||
def forward(self, sigclip_embeds) -> torch.Tensor:
|
||||
x = self.redux_up(sigclip_embeds)
|
||||
x = torch.nn.functional.silu(x)
|
||||
|
||||
projected_x = self.redux_down(x)
|
||||
return projected_x
|
||||
61
toolkit/models/sref.py
Normal file
61
toolkit/models/sref.py
Normal file
@@ -0,0 +1,61 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class SrefImageEncoder(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_features: int = 1152,
|
||||
input_tokens: int = 512,
|
||||
output_tokens: int = 512,
|
||||
output_features: int = 4096,
|
||||
intermediate_size: int = 4096,
|
||||
num_digits: int = 10,
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.input_features = input_features
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.input_tokens = input_tokens
|
||||
self.output_tokens = output_tokens
|
||||
self.output_features = output_features
|
||||
self.intermediate_size = intermediate_size
|
||||
self.num_digits = num_digits
|
||||
|
||||
self.proj_in = nn.Linear(
|
||||
input_features, intermediate_size, dtype=dtype)
|
||||
# (bs, num_digits, intermediate_size)
|
||||
self.conv_pool = nn.Conv1d(input_tokens, num_digits, 1, dtype=dtype)
|
||||
self.linear_pool = nn.Linear(
|
||||
intermediate_size, 1, dtype=dtype) # (bs, num_digits, 1)
|
||||
# do sigmoid for digits 0.0-1.0 = (0 to 10) Always floor when rounding digits so you get 0-9
|
||||
self.flatten = nn.Flatten() # (bs, num_digits * intermediate_size)
|
||||
|
||||
# a numeric sref would come in here with num_digits
|
||||
self.sref_in = nn.Linear(num_digits, intermediate_size, dtype=dtype)
|
||||
self.fc1 = nn.Linear(intermediate_size, intermediate_size, dtype=dtype)
|
||||
self.fc2 = nn.Linear(intermediate_size, intermediate_size, dtype=dtype)
|
||||
|
||||
self.proj_out = nn.Linear(
|
||||
intermediate_size, output_features * output_tokens, dtype=dtype)
|
||||
|
||||
def forward(self, siglip_embeds) -> torch.Tensor:
|
||||
x = self.proj_in(siglip_embeds)
|
||||
x = torch.nn.functional.silu(x)
|
||||
x = self.conv_pool(x)
|
||||
x = self.linear_pool(x)
|
||||
x = torch.sigmoid(x)
|
||||
|
||||
sref = self.flatten(x)
|
||||
|
||||
x = self.sref_in(sref)
|
||||
x = torch.nn.functional.silu(x)
|
||||
x = self.fc1(x)
|
||||
x = torch.nn.functional.silu(x)
|
||||
x = self.fc2(x)
|
||||
x = torch.nn.functional.silu(x)
|
||||
x = self.proj_out(x)
|
||||
|
||||
return x
|
||||
@@ -5,9 +5,14 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import weakref
|
||||
from typing import Union, TYPE_CHECKING, Optional
|
||||
from collections import OrderedDict
|
||||
|
||||
from diffusers import Transformer2DModel, FluxTransformer2DModel
|
||||
from transformers import T5EncoderModel, CLIPTextModel, CLIPTokenizer, T5Tokenizer, CLIPVisionModelWithProjection
|
||||
from toolkit.models.pixtral_vision import PixtralVisionEncoder, PixtralVisionImagePreprocessor, VisionLanguageAdapter
|
||||
from transformers import SiglipImageProcessor, SiglipVisionModel
|
||||
|
||||
from toolkit.config_modules import AdapterConfig
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
sys.path.append(REPOS_ROOT)
|
||||
|
||||
@@ -15,29 +20,81 @@ sys.path.append(REPOS_ROOT)
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
|
||||
|
||||
# matches distribution of randn
|
||||
class Norm(nn.Module):
|
||||
def __init__(self, target_mean=0.0, target_std=1.0, eps=1e-6):
|
||||
super(Norm, self).__init__()
|
||||
self.target_mean = target_mean
|
||||
self.target_std = target_std
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x):
|
||||
dims = tuple(range(1, x.dim()))
|
||||
mean = x.mean(dim=dims, keepdim=True)
|
||||
std = x.std(dim=dims, keepdim=True)
|
||||
|
||||
# Normalize
|
||||
return self.target_std * (x - mean) / (std + self.eps) + self.target_mean
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
|
||||
norm_layer = Norm()
|
||||
|
||||
class SparseAutoencoder(nn.Module):
|
||||
def __init__(self, input_dim, hidden_dim, output_dim):
|
||||
super(SparseAutoencoder, self).__init__()
|
||||
self.encoder = nn.Sequential(
|
||||
nn.Linear(input_dim, hidden_dim),
|
||||
nn.GELU(),
|
||||
nn.Linear(hidden_dim, output_dim),
|
||||
)
|
||||
self.norm = Norm()
|
||||
self.decoder = nn.Sequential(
|
||||
nn.Linear(output_dim, hidden_dim),
|
||||
nn.GELU(),
|
||||
nn.Linear(hidden_dim, input_dim),
|
||||
)
|
||||
self.last_run = None
|
||||
|
||||
def forward(self, x):
|
||||
self.last_run = {
|
||||
"input": x
|
||||
}
|
||||
x = self.encoder(x)
|
||||
x = self.norm(x)
|
||||
self.last_run["sparse"] = x
|
||||
x = self.decoder(x)
|
||||
x = self.norm(x)
|
||||
self.last_run["output"] = x
|
||||
return x
|
||||
|
||||
|
||||
class MLPR(nn.Module): # MLP with reshaping
|
||||
def __init__(
|
||||
self,
|
||||
in_dim,
|
||||
in_channels,
|
||||
out_dim,
|
||||
out_channels,
|
||||
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
|
||||
# dont normalize if using conv
|
||||
self.layer_norm = nn.LayerNorm(in_dim)
|
||||
|
||||
self.fc1 = nn.Linear(in_dim, out_dim)
|
||||
self.act_fn = nn.GELU()
|
||||
self.conv1 = nn.Conv1d(in_channels, out_channels, 1)
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.layer_norm(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
|
||||
x = self.conv1(x)
|
||||
return x
|
||||
|
||||
class AttnProcessor2_0(torch.nn.Module):
|
||||
@@ -286,7 +343,7 @@ 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):
|
||||
adapter_hidden_size=None, has_bias=False, block_idx=0, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
@@ -298,6 +355,7 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
self.adapter_hidden_size = adapter_hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
self.block_idx = block_idx
|
||||
|
||||
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)
|
||||
@@ -323,17 +381,7 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
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]
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
@@ -352,36 +400,34 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
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)
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `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)
|
||||
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)
|
||||
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)
|
||||
# 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)
|
||||
@@ -391,10 +437,14 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
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
|
||||
# begin ip adapter
|
||||
if self.is_active and self.conditional_embeds is not None:
|
||||
adapter_hidden_states = self.conditional_embeds
|
||||
block_scaler = self.adapter_ref().block_scaler
|
||||
if block_scaler is not None:
|
||||
# add 1 to block scaler so we can decay its weight to 1.0
|
||||
block_scaler = block_scaler[self.block_idx] + 1.0
|
||||
|
||||
if adapter_hidden_states.shape[0] < batch_size:
|
||||
adapter_hidden_states = torch.cat([
|
||||
self.unconditional_embeds,
|
||||
@@ -413,8 +463,6 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
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
|
||||
)
|
||||
@@ -422,27 +470,32 @@ class CustomFluxVDAttnProcessor2_0(torch.nn.Module):
|
||||
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)
|
||||
|
||||
# scale to block scaler
|
||||
if block_scaler is not None:
|
||||
orig_dtype = vd_hidden_states.dtype
|
||||
if block_scaler.dtype != vd_hidden_states.dtype:
|
||||
vd_hidden_states = vd_hidden_states.to(block_scaler.dtype)
|
||||
vd_hidden_states = vd_hidden_states * block_scaler
|
||||
if block_scaler.dtype != orig_dtype:
|
||||
vd_hidden_states = vd_hidden_states.to(orig_dtype)
|
||||
|
||||
hidden_states = hidden_states + self.scale * vd_hidden_states
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
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
|
||||
# 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)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
|
||||
class VisionDirectAdapter(torch.nn.Module):
|
||||
def __init__(
|
||||
@@ -456,12 +509,31 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
is_flux = sd.is_flux
|
||||
self.adapter_ref: weakref.ref = weakref.ref(adapter)
|
||||
self.sd_ref: weakref.ref = weakref.ref(sd)
|
||||
self.config: AdapterConfig = adapter.config
|
||||
self.vision_model_ref: weakref.ref = weakref.ref(vision_model)
|
||||
self.resampler = None
|
||||
is_pixtral = self.config.image_encoder_arch == "pixtral"
|
||||
|
||||
if adapter.config.clip_layer == "image_embeds":
|
||||
self.token_size = vision_model.config.projection_dim
|
||||
if isinstance(vision_model, SiglipVisionModel):
|
||||
self.token_size = vision_model.config.hidden_size
|
||||
else:
|
||||
self.token_size = vision_model.config.projection_dim
|
||||
else:
|
||||
self.token_size = vision_model.config.hidden_size
|
||||
|
||||
self.mid_size = self.token_size
|
||||
|
||||
if self.config.conv_pooling and self.config.conv_pooling_stacks > 1:
|
||||
self.mid_size = self.mid_size * self.config.conv_pooling_stacks
|
||||
|
||||
# if pixtral, use cross attn dim for more sparse representation if only doing double transformers
|
||||
if is_pixtral and self.config.flux_only_double:
|
||||
if is_flux:
|
||||
hidden_size = 3072
|
||||
else:
|
||||
hidden_size = sd.unet.config['cross_attention_dim']
|
||||
self.mid_size = hidden_size
|
||||
|
||||
# init adapter modules
|
||||
attn_procs = {}
|
||||
@@ -482,12 +554,15 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
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")
|
||||
if not self.config.flux_only_double:
|
||||
# single transformer blocks do not have cross attn, but we will do them anyway
|
||||
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())
|
||||
|
||||
current_idx = 0
|
||||
|
||||
for name in attn_processor_keys:
|
||||
if is_flux:
|
||||
cross_attention_dim = None
|
||||
@@ -501,7 +576,7 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
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"):
|
||||
elif name.startswith("transformer") or name.startswith("single_transformer"):
|
||||
if is_flux:
|
||||
hidden_size = 3072
|
||||
else:
|
||||
@@ -525,27 +600,27 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
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:
|
||||
if to_k_adapter.shape[1] < self.mid_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(
|
||||
torch.randn(to_k_adapter.shape[0], self.mid_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(
|
||||
torch.randn(to_v_adapter.shape[0], self.mid_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]
|
||||
elif to_k_adapter.shape[1] > self.mid_size:
|
||||
to_k_adapter = to_k_adapter[:, :self.mid_size]
|
||||
to_v_adapter = to_v_adapter[:, :self.mid_size]
|
||||
# if is_pixart:
|
||||
# to_k_bias = to_k_bias[:self.token_size]
|
||||
# to_v_bias = to_v_bias[:self.token_size]
|
||||
# to_k_bias = to_k_bias[:self.mid_size]
|
||||
# to_v_bias = to_v_bias[:self.mid_size]
|
||||
else:
|
||||
to_k_adapter = to_k_adapter
|
||||
to_v_adapter = to_v_adapter
|
||||
@@ -567,8 +642,9 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
adapter_hidden_size=self.mid_size,
|
||||
has_bias=False,
|
||||
block_idx=current_idx
|
||||
)
|
||||
else:
|
||||
attn_procs[name] = VisionDirectAdapterAttnProcessor(
|
||||
@@ -576,10 +652,12 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
scale=1.0,
|
||||
adapter=self,
|
||||
adapter_hidden_size=self.token_size,
|
||||
adapter_hidden_size=self.mid_size,
|
||||
has_bias=False,
|
||||
)
|
||||
current_idx += 1
|
||||
attn_procs[name].load_state_dict(weights)
|
||||
|
||||
if self.sd_ref().is_pixart:
|
||||
# we have to set them ourselves
|
||||
transformer: Transformer2DModel = sd.unet
|
||||
@@ -596,23 +674,106 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
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))
|
||||
])
|
||||
|
||||
if not self.config.flux_only_double:
|
||||
# do single blocks too even though they dont have cross attn
|
||||
for i, module in transformer.single_transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"single_transformer_blocks.{i}.attn"]
|
||||
|
||||
if not self.config.flux_only_double:
|
||||
self.adapter_modules = torch.nn.ModuleList(
|
||||
[
|
||||
transformer.transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.transformer_blocks))
|
||||
] + [
|
||||
transformer.single_transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.single_transformer_blocks))
|
||||
]
|
||||
)
|
||||
else:
|
||||
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
|
||||
)
|
||||
num_modules = len(self.adapter_modules)
|
||||
if self.config.train_scaler:
|
||||
self.block_scaler = torch.nn.Parameter(torch.tensor([0.0] * num_modules).to(
|
||||
dtype=torch.float32,
|
||||
device=self.sd_ref().device_torch
|
||||
))
|
||||
self.block_scaler.data = self.block_scaler.data.to(torch.float32)
|
||||
self.block_scaler.requires_grad = True
|
||||
else:
|
||||
self.block_scaler = None
|
||||
|
||||
self.pool = None
|
||||
|
||||
if self.config.num_tokens is not None:
|
||||
# image_encoder_state_dict = self.adapter_ref().vision_encoder.state_dict()
|
||||
# max_seq_len = CLIP tokens + CLS token
|
||||
# max_seq_len = 257
|
||||
# if "vision_model.embeddings.position_embedding.weight" in image_encoder_state_dict:
|
||||
# # clip
|
||||
# max_seq_len = int(
|
||||
# image_encoder_state_dict["vision_model.embeddings.position_embedding.weight"].shape[0])
|
||||
# self.resampler = MLPR(
|
||||
# in_dim=self.token_size,
|
||||
# in_channels=max_seq_len,
|
||||
# out_dim=self.mid_size,
|
||||
# out_channels=self.config.num_tokens,
|
||||
# )
|
||||
vision_config = self.adapter_ref().vision_encoder.config
|
||||
# sequence_length = int((vision_config.image_size / vision_config.patch_size) ** 2 + 1)
|
||||
# siglip doesnt add 1
|
||||
sequence_length = int((vision_config.image_size / vision_config.patch_size) ** 2)
|
||||
self.pool = nn.Sequential(
|
||||
nn.Conv1d(sequence_length, self.config.num_tokens, 1, bias=False),
|
||||
Norm(),
|
||||
)
|
||||
|
||||
elif self.config.image_encoder_arch == "pixtral":
|
||||
self.resampler = VisionLanguageAdapter(
|
||||
in_dim=self.token_size,
|
||||
out_dim=self.mid_size,
|
||||
)
|
||||
|
||||
self.sparse_autoencoder = None
|
||||
if self.config.conv_pooling:
|
||||
vision_config = self.adapter_ref().vision_encoder.config
|
||||
# sequence_length = int((vision_config.image_size / vision_config.patch_size) ** 2 + 1)
|
||||
# siglip doesnt add 1
|
||||
sequence_length = int((vision_config.image_size / vision_config.patch_size) ** 2)
|
||||
self.pool = nn.Sequential(
|
||||
nn.Conv1d(sequence_length, self.config.conv_pooling_stacks, 1, bias=False),
|
||||
Norm(),
|
||||
)
|
||||
if self.config.sparse_autoencoder_dim is not None:
|
||||
hidden_dim = self.token_size * 2
|
||||
if hidden_dim > self.config.sparse_autoencoder_dim:
|
||||
hidden_dim = self.config.sparse_autoencoder_dim
|
||||
self.sparse_autoencoder = SparseAutoencoder(
|
||||
input_dim=self.token_size,
|
||||
hidden_dim=hidden_dim,
|
||||
output_dim=self.config.sparse_autoencoder_dim
|
||||
)
|
||||
|
||||
if self.config.clip_layer == "image_embeds":
|
||||
self.proj = nn.Linear(self.token_size, self.token_size)
|
||||
|
||||
def state_dict(self, destination=None, prefix='', keep_vars=False):
|
||||
if self.config.train_scaler:
|
||||
# only return the block scaler
|
||||
if destination is None:
|
||||
destination = OrderedDict()
|
||||
destination[prefix + 'block_scaler'] = self.block_scaler
|
||||
return destination
|
||||
return super().state_dict(destination, prefix, keep_vars)
|
||||
|
||||
# make a getter to see if is active
|
||||
@property
|
||||
@@ -620,4 +781,32 @@ class VisionDirectAdapter(torch.nn.Module):
|
||||
return self.adapter_ref().is_active
|
||||
|
||||
def forward(self, input):
|
||||
return self.mlp(input)
|
||||
# block scaler keeps moving dtypes. make sure it is float32 here
|
||||
# todo remove this when we have a real solution
|
||||
|
||||
if self.block_scaler is not None and self.block_scaler.dtype != torch.float32:
|
||||
self.block_scaler.data = self.block_scaler.data.to(torch.float32)
|
||||
# if doing image_embeds, normalize here
|
||||
if self.config.clip_layer == "image_embeds":
|
||||
input = norm_layer(input)
|
||||
input = self.proj(input)
|
||||
if self.resampler is not None:
|
||||
input = self.resampler(input)
|
||||
if self.pool is not None:
|
||||
input = self.pool(input)
|
||||
if self.config.conv_pooling_stacks > 1:
|
||||
input = torch.cat(torch.chunk(input, self.config.conv_pooling_stacks, dim=1), dim=2)
|
||||
if self.sparse_autoencoder is not None:
|
||||
input = self.sparse_autoencoder(input)
|
||||
return input
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
super().to(*args, **kwargs)
|
||||
if self.block_scaler is not None:
|
||||
if self.block_scaler.dtype != torch.float32:
|
||||
self.block_scaler.data = self.block_scaler.data.to(torch.float32)
|
||||
return self
|
||||
|
||||
def post_weight_update(self):
|
||||
# force block scaler to be mean of 1
|
||||
pass
|
||||
|
||||
1
toolkit/models/wan21/__init__.py
Normal file
1
toolkit/models/wan21/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
from .wan21 import Wan21
|
||||
679
toolkit/models/wan21/wan21.py
Normal file
679
toolkit/models/wan21/wan21.py
Normal file
@@ -0,0 +1,679 @@
|
||||
# WIP, coming soon ish
|
||||
from functools import partial
|
||||
import torch
|
||||
import yaml
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from toolkit.basic import flush
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
from toolkit.paths import REPOS_ROOT
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
from diffusers import AutoencoderKLWan, WanPipeline, WanTransformer3DModel
|
||||
import os
|
||||
import sys
|
||||
|
||||
import weakref
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from toolkit.basic import flush
|
||||
from toolkit.config_modules import GenerateImageConfig, ModelConfig
|
||||
from toolkit.dequantize import patch_dequantization_on_save
|
||||
from toolkit.models.base_model import BaseModel
|
||||
from toolkit.prompt_utils import PromptEmbeds
|
||||
|
||||
import os
|
||||
import copy
|
||||
from toolkit.config_modules import ModelConfig, GenerateImageConfig, ModelArch
|
||||
import torch
|
||||
from optimum.quanto import freeze, qfloat8, QTensor, qint4
|
||||
from toolkit.util.quantize import quantize, get_qtype
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler, UniPCMultistepScheduler
|
||||
from typing import TYPE_CHECKING, List
|
||||
from toolkit.accelerator import unwrap_model
|
||||
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
|
||||
from torchvision.transforms import Resize, ToPILImage
|
||||
from tqdm import tqdm
|
||||
|
||||
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
|
||||
from diffusers.pipelines.wan.pipeline_wan import XLA_AVAILABLE
|
||||
# from ...callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
from toolkit.models.wan21.wan_lora_convert import convert_to_diffusers, convert_to_original
|
||||
|
||||
# for generation only?
|
||||
scheduler_configUniPC = {
|
||||
"_class_name": "UniPCMultistepScheduler",
|
||||
"_diffusers_version": "0.33.0.dev0",
|
||||
"beta_end": 0.02,
|
||||
"beta_schedule": "linear",
|
||||
"beta_start": 0.0001,
|
||||
"disable_corrector": [],
|
||||
"dynamic_thresholding_ratio": 0.995,
|
||||
"final_sigmas_type": "zero",
|
||||
"flow_shift": 3.0,
|
||||
"lower_order_final": True,
|
||||
"num_train_timesteps": 1000,
|
||||
"predict_x0": True,
|
||||
"prediction_type": "flow_prediction",
|
||||
"rescale_betas_zero_snr": False,
|
||||
"sample_max_value": 1.0,
|
||||
"solver_order": 2,
|
||||
"solver_p": None,
|
||||
"solver_type": "bh2",
|
||||
"steps_offset": 0,
|
||||
"thresholding": False,
|
||||
"timestep_spacing": "linspace",
|
||||
"trained_betas": None,
|
||||
"use_beta_sigmas": False,
|
||||
"use_exponential_sigmas": False,
|
||||
"use_flow_sigmas": True,
|
||||
"use_karras_sigmas": False
|
||||
}
|
||||
|
||||
# for training. I think it is right
|
||||
scheduler_config = {
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": False
|
||||
}
|
||||
|
||||
|
||||
class AggressiveWanUnloadPipeline(WanPipeline):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
text_encoder: UMT5EncoderModel,
|
||||
transformer: WanTransformer3DModel,
|
||||
vae: AutoencoderKLWan,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
device: torch.device = torch.device("cuda"),
|
||||
):
|
||||
super().__init__(
|
||||
tokenizer=tokenizer,
|
||||
text_encoder=text_encoder,
|
||||
transformer=transformer,
|
||||
vae=vae,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
self._exec_device = device
|
||||
@property
|
||||
def _execution_device(self):
|
||||
return self._exec_device
|
||||
|
||||
def __call__(
|
||||
self: WanPipeline,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator,
|
||||
List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "np",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[
|
||||
Union[Callable[[int, int, Dict], None],
|
||||
PipelineCallback, MultiPipelineCallbacks]
|
||||
] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# unload vae and transformer
|
||||
vae_device = self.vae.device
|
||||
transformer_device = self.transformer.device
|
||||
text_encoder_device = self.text_encoder.device
|
||||
device = self.transformer.device
|
||||
|
||||
print("Unloading vae")
|
||||
self.vae.to("cpu")
|
||||
self.text_encoder.to(device)
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._current_timestep = None
|
||||
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]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# unload text encoder
|
||||
print("Unloading text encoder")
|
||||
self.text_encoder.to("cpu")
|
||||
|
||||
self.transformer.to(device)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(device, transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(
|
||||
device, transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 6. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - \
|
||||
num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
latent_model_input = latents.to(device, transformer_dtype)
|
||||
timestep = t.expand(latents.shape[0])
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_uncond = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * \
|
||||
(noise_pred - noise_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(
|
||||
noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
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)
|
||||
negative_prompt_embeds = callback_outputs.pop(
|
||||
"negative_prompt_embeds", negative_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()
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
# unload transformer
|
||||
# load vae
|
||||
print("Loading Vae")
|
||||
self.vae.to(vae_device)
|
||||
|
||||
if not output_type == "latent":
|
||||
latents = latents.to(self.vae.dtype)
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean)
|
||||
.view(1, self.vae.config.z_dim, 1, 1, 1)
|
||||
.to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = latents / latents_std + latents_mean
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(
|
||||
video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return WanPipelineOutput(frames=video)
|
||||
|
||||
|
||||
class Wan21(BaseModel):
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
model_config: ModelConfig,
|
||||
dtype='bf16',
|
||||
custom_pipeline=None,
|
||||
noise_scheduler=None,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__(device, model_config, dtype,
|
||||
custom_pipeline, noise_scheduler, **kwargs)
|
||||
self.is_flow_matching = True
|
||||
self.is_transformer = True
|
||||
self.target_lora_modules = ['WanTransformer3DModel']
|
||||
|
||||
# cache for holding noise
|
||||
self.effective_noise = None
|
||||
|
||||
# static method to get the scheduler
|
||||
@staticmethod
|
||||
def get_train_scheduler():
|
||||
scheduler = CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
|
||||
return scheduler
|
||||
|
||||
def load_model(self):
|
||||
dtype = self.torch_dtype
|
||||
# todo , will this work with other wan models?
|
||||
base_model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
model_path = self.model_config.name_or_path
|
||||
|
||||
self.print_and_status_update("Loading Wan2.1 model")
|
||||
# base_model_path = "black-forest-labs/FLUX.1-schnell"
|
||||
base_model_path = self.model_config.name_or_path_original
|
||||
subfolder = 'transformer'
|
||||
transformer_path = model_path
|
||||
if os.path.exists(transformer_path):
|
||||
subfolder = None
|
||||
transformer_path = os.path.join(transformer_path, 'transformer')
|
||||
# check if the path is a full checkpoint.
|
||||
te_folder_path = os.path.join(model_path, 'text_encoder')
|
||||
# if we have the te, this folder is a full checkpoint, use it as the base
|
||||
if os.path.exists(te_folder_path):
|
||||
base_model_path = model_path
|
||||
|
||||
self.print_and_status_update("Loading transformer")
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
transformer_path,
|
||||
subfolder=subfolder,
|
||||
torch_dtype=dtype,
|
||||
).to(dtype=dtype)
|
||||
|
||||
if self.model_config.split_model_over_gpus:
|
||||
raise ValueError(
|
||||
"Splitting model over gpus is not supported for Wan2.1 models")
|
||||
|
||||
if not self.model_config.low_vram:
|
||||
# quantize on the device
|
||||
transformer.to(self.quantize_device, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.assistant_lora_path is not None or self.model_config.inference_lora_path is not None:
|
||||
raise ValueError(
|
||||
"Assistant LoRA is not supported for Wan2.1 models currently")
|
||||
|
||||
if self.model_config.lora_path is not None:
|
||||
raise ValueError(
|
||||
"Loading LoRA is not supported for Wan2.1 models currently")
|
||||
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize:
|
||||
print("Quantizing Transformer")
|
||||
quantization_args = self.model_config.quantize_kwargs
|
||||
if 'exclude' not in quantization_args:
|
||||
quantization_args['exclude'] = []
|
||||
# patch the state dict method
|
||||
patch_dequantization_on_save(transformer)
|
||||
quantization_type = get_qtype(self.model_config.qtype)
|
||||
self.print_and_status_update("Quantizing transformer")
|
||||
if self.model_config.low_vram:
|
||||
print("Quantizing blocks")
|
||||
orig_exclude = copy.deepcopy(quantization_args['exclude'])
|
||||
# quantize each block
|
||||
idx = 0
|
||||
for block in tqdm(transformer.blocks):
|
||||
block.to(self.device_torch)
|
||||
quantize(block, weights=quantization_type,
|
||||
**quantization_args)
|
||||
freeze(block)
|
||||
idx += 1
|
||||
flush()
|
||||
|
||||
print("Quantizing the rest")
|
||||
low_vram_exclude = copy.deepcopy(quantization_args['exclude'])
|
||||
low_vram_exclude.append('blocks.*')
|
||||
quantization_args['exclude'] = low_vram_exclude
|
||||
# quantize the rest
|
||||
transformer.to(self.device_torch)
|
||||
quantize(transformer, weights=quantization_type,
|
||||
**quantization_args)
|
||||
|
||||
quantization_args['exclude'] = orig_exclude
|
||||
else:
|
||||
# do it in one go
|
||||
quantize(transformer, weights=quantization_type,
|
||||
**quantization_args)
|
||||
freeze(transformer)
|
||||
# move it to the cpu for now
|
||||
transformer.to("cpu")
|
||||
else:
|
||||
transformer.to(self.device_torch, dtype=dtype)
|
||||
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Loading UMT5EncoderModel")
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
base_model_path, subfolder="tokenizer", torch_dtype=dtype)
|
||||
text_encoder = UMT5EncoderModel.from_pretrained(
|
||||
base_model_path, subfolder="text_encoder", torch_dtype=dtype).to(dtype=dtype)
|
||||
|
||||
text_encoder.to(self.device_torch, dtype=dtype)
|
||||
flush()
|
||||
|
||||
if self.model_config.quantize_te:
|
||||
self.print_and_status_update("Quantizing UMT5EncoderModel")
|
||||
quantize(text_encoder, weights=get_qtype(self.model_config.qtype))
|
||||
freeze(text_encoder)
|
||||
flush()
|
||||
|
||||
if self.model_config.low_vram:
|
||||
print("Moving transformer back to GPU")
|
||||
# we can move it back to the gpu now
|
||||
transformer.to(self.device_torch)
|
||||
|
||||
scheduler = Wan21.get_train_scheduler()
|
||||
self.print_and_status_update("Loading VAE")
|
||||
# todo, example does float 32? check if quality suffers
|
||||
vae = AutoencoderKLWan.from_pretrained(
|
||||
base_model_path, subfolder="vae", torch_dtype=dtype).to(dtype=dtype)
|
||||
flush()
|
||||
|
||||
self.print_and_status_update("Making pipe")
|
||||
pipe: WanPipeline = WanPipeline(
|
||||
scheduler=scheduler,
|
||||
text_encoder=None,
|
||||
tokenizer=tokenizer,
|
||||
vae=vae,
|
||||
transformer=None,
|
||||
)
|
||||
pipe.text_encoder = text_encoder
|
||||
pipe.transformer = transformer
|
||||
|
||||
self.print_and_status_update("Preparing Model")
|
||||
|
||||
text_encoder = pipe.text_encoder
|
||||
tokenizer = pipe.tokenizer
|
||||
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
|
||||
flush()
|
||||
text_encoder.to(self.device_torch)
|
||||
text_encoder.requires_grad_(False)
|
||||
text_encoder.eval()
|
||||
pipe.transformer = pipe.transformer.to(self.device_torch)
|
||||
flush()
|
||||
self.pipeline = pipe
|
||||
self.model = transformer
|
||||
self.vae = vae
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def get_generation_pipeline(self):
|
||||
scheduler = UniPCMultistepScheduler(**scheduler_configUniPC)
|
||||
if self.model_config.low_vram:
|
||||
pipeline = AggressiveWanUnloadPipeline(
|
||||
vae=self.vae,
|
||||
transformer=self.model,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=scheduler,
|
||||
device=self.device_torch
|
||||
)
|
||||
else:
|
||||
pipeline = WanPipeline(
|
||||
vae=self.vae,
|
||||
transformer=self.unet,
|
||||
text_encoder=self.text_encoder,
|
||||
tokenizer=self.tokenizer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
|
||||
return pipeline
|
||||
|
||||
def generate_single_image(
|
||||
self,
|
||||
pipeline: WanPipeline,
|
||||
gen_config: GenerateImageConfig,
|
||||
conditional_embeds: PromptEmbeds,
|
||||
unconditional_embeds: PromptEmbeds,
|
||||
generator: torch.Generator,
|
||||
extra: dict,
|
||||
):
|
||||
# reactivate progress bar since this is slooooow
|
||||
pipeline.set_progress_bar_config(disable=False)
|
||||
pipeline = pipeline.to(self.device_torch)
|
||||
# todo, figure out how to do video
|
||||
output = pipeline(
|
||||
prompt_embeds=conditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
negative_prompt_embeds=unconditional_embeds.text_embeds.to(
|
||||
self.device_torch, dtype=self.torch_dtype),
|
||||
height=gen_config.height,
|
||||
width=gen_config.width,
|
||||
num_inference_steps=gen_config.num_inference_steps,
|
||||
guidance_scale=gen_config.guidance_scale,
|
||||
latents=gen_config.latents,
|
||||
num_frames=gen_config.num_frames,
|
||||
generator=generator,
|
||||
return_dict=False,
|
||||
output_type="pil",
|
||||
**extra
|
||||
)[0]
|
||||
|
||||
# shape = [1, frames, channels, height, width]
|
||||
batch_item = output[0] # list of pil images
|
||||
if gen_config.num_frames > 1:
|
||||
return batch_item # return the frames.
|
||||
else:
|
||||
# get just the first image
|
||||
img = batch_item[0]
|
||||
return img
|
||||
|
||||
def get_noise_prediction(
|
||||
self,
|
||||
latent_model_input: torch.Tensor,
|
||||
timestep: torch.Tensor, # 0 to 1000 scale
|
||||
text_embeddings: PromptEmbeds,
|
||||
**kwargs
|
||||
):
|
||||
# vae_scale_factor_spatial = 8
|
||||
# vae_scale_factor_temporal = 4
|
||||
# num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
# shape = (
|
||||
# batch_size,
|
||||
# num_channels_latents, # 16
|
||||
# num_latent_frames, # 81
|
||||
# int(height) // self.vae_scale_factor_spatial,
|
||||
# int(width) // self.vae_scale_factor_spatial,
|
||||
# )
|
||||
|
||||
noise_pred = self.model(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=text_embeddings.text_embeds,
|
||||
return_dict=False,
|
||||
**kwargs
|
||||
)[0]
|
||||
return noise_pred
|
||||
|
||||
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
|
||||
if self.pipeline.text_encoder.device != self.device_torch:
|
||||
self.pipeline.text_encoder.to(self.device_torch)
|
||||
prompt_embeds, _ = self.pipeline.encode_prompt(
|
||||
prompt,
|
||||
do_classifier_free_guidance=False,
|
||||
max_sequence_length=512,
|
||||
device=self.device_torch,
|
||||
dtype=self.torch_dtype,
|
||||
)
|
||||
return PromptEmbeds(prompt_embeds)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_images(
|
||||
self,
|
||||
image_list: List[torch.Tensor],
|
||||
device=None,
|
||||
dtype=None
|
||||
):
|
||||
if device is None:
|
||||
device = self.vae_device_torch
|
||||
if dtype is None:
|
||||
dtype = self.vae_torch_dtype
|
||||
|
||||
# Move to vae to device if on cpu
|
||||
if self.vae.device == 'cpu':
|
||||
self.vae.to(device)
|
||||
self.vae.eval()
|
||||
self.vae.requires_grad_(False)
|
||||
# move to device and dtype
|
||||
image_list = [image.to(device, dtype=dtype) for image in image_list]
|
||||
|
||||
# We need to detect video if we have it.
|
||||
# videos come in (num_frames, channels, height, width)
|
||||
# images come in (channels, height, width)
|
||||
# we need to add a frame dimension to images and remap the video to (channels, num_frames, height, width)
|
||||
|
||||
if len(image_list[0].shape) == 3:
|
||||
image_list = [image.unsqueeze(1) for image in image_list]
|
||||
elif len(image_list[0].shape) == 4:
|
||||
image_list = [image.permute(1, 0, 2, 3) for image in image_list]
|
||||
else:
|
||||
raise ValueError(f"Image shape is not correct, got {list(image_list[0].shape)}")
|
||||
|
||||
VAE_SCALE_FACTOR = 8
|
||||
|
||||
# resize images if not divisible by 8
|
||||
# now we need to resize considering the shape (channels, num_frames, height, width)
|
||||
for i in range(len(image_list)):
|
||||
image = image_list[i]
|
||||
if image.shape[2] % VAE_SCALE_FACTOR != 0 or image.shape[3] % VAE_SCALE_FACTOR != 0:
|
||||
# Create resized frames by handling each frame separately
|
||||
c, f, h, w = image.shape
|
||||
target_h = h // VAE_SCALE_FACTOR * VAE_SCALE_FACTOR
|
||||
target_w = w // VAE_SCALE_FACTOR * VAE_SCALE_FACTOR
|
||||
|
||||
# We need to process each frame separately
|
||||
resized_frames = []
|
||||
for frame_idx in range(f):
|
||||
frame = image[:, frame_idx, :, :] # Extract single frame (channels, height, width)
|
||||
resized_frame = Resize((target_h, target_w))(frame)
|
||||
resized_frames.append(resized_frame.unsqueeze(1)) # Add frame dimension back
|
||||
|
||||
# Concatenate all frames back together along the frame dimension
|
||||
image_list[i] = torch.cat(resized_frames, dim=1)
|
||||
|
||||
images = torch.stack(image_list)
|
||||
# images = images.unsqueeze(2) # adds frame dimension so (bs, ch, h, w) -> (bs, ch, 1, h, w)
|
||||
latents = self.vae.encode(images).latent_dist.sample()
|
||||
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean)
|
||||
.view(1, self.vae.config.z_dim, 1, 1, 1)
|
||||
.to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = (latents - latents_mean) * latents_std
|
||||
|
||||
latents = latents.to(device, dtype=dtype)
|
||||
|
||||
return latents
|
||||
|
||||
def get_model_has_grad(self):
|
||||
return self.model.proj_out.weight.requires_grad
|
||||
|
||||
def get_te_has_grad(self):
|
||||
return self.text_encoder.encoder.block[0].layer[0].SelfAttention.q.weight.requires_grad
|
||||
|
||||
def save_model(self, output_path, meta, save_dtype):
|
||||
# only save the unet
|
||||
transformer: Wan21 = unwrap_model(self.model)
|
||||
transformer.save_pretrained(
|
||||
save_directory=os.path.join(output_path, 'transformer'),
|
||||
safe_serialization=True,
|
||||
)
|
||||
|
||||
meta_path = os.path.join(output_path, 'aitk_meta.yaml')
|
||||
with open(meta_path, 'w') as f:
|
||||
yaml.dump(meta, f)
|
||||
|
||||
def get_loss_target(self, *args, **kwargs):
|
||||
noise = kwargs.get('noise')
|
||||
batch = kwargs.get('batch')
|
||||
if batch is None:
|
||||
raise ValueError("Batch is not provided")
|
||||
if noise is None:
|
||||
raise ValueError("Noise is not provided")
|
||||
return (noise - batch.latents).detach()
|
||||
|
||||
def convert_lora_weights_before_save(self, state_dict):
|
||||
return convert_to_original(state_dict)
|
||||
|
||||
def convert_lora_weights_before_load(self, state_dict):
|
||||
return convert_to_diffusers(state_dict)
|
||||
65
toolkit/models/wan21/wan_lora_convert.py
Normal file
65
toolkit/models/wan21/wan_lora_convert.py
Normal file
@@ -0,0 +1,65 @@
|
||||
def convert_to_diffusers(state_dict):
|
||||
new_state_dict = {}
|
||||
for key in state_dict:
|
||||
new_key = key
|
||||
# Base model name change
|
||||
if key.startswith("diffusion_model."):
|
||||
new_key = key.replace("diffusion_model.", "transformer.")
|
||||
|
||||
# Attention blocks conversion
|
||||
if "self_attn" in new_key:
|
||||
new_key = new_key.replace("self_attn", "attn1")
|
||||
elif "cross_attn" in new_key:
|
||||
new_key = new_key.replace("cross_attn", "attn2")
|
||||
|
||||
# Attention components conversion
|
||||
parts = new_key.split(".")
|
||||
for i, part in enumerate(parts):
|
||||
if part in ["q", "k", "v"]:
|
||||
parts[i] = f"to_{part}"
|
||||
elif part == "o":
|
||||
parts[i] = "to_out.0"
|
||||
new_key = ".".join(parts)
|
||||
|
||||
# FFN conversion
|
||||
if "ffn.0" in new_key:
|
||||
new_key = new_key.replace("ffn.0", "ffn.net.0.proj")
|
||||
elif "ffn.2" in new_key:
|
||||
new_key = new_key.replace("ffn.2", "ffn.net.2")
|
||||
|
||||
new_state_dict[new_key] = state_dict[key]
|
||||
return new_state_dict
|
||||
|
||||
|
||||
def convert_to_original(state_dict):
|
||||
new_state_dict = {}
|
||||
for key in state_dict:
|
||||
new_key = key
|
||||
# Base model name change
|
||||
if key.startswith("transformer."):
|
||||
new_key = key.replace("transformer.", "diffusion_model.")
|
||||
|
||||
# Attention blocks conversion
|
||||
if "attn1" in new_key:
|
||||
new_key = new_key.replace("attn1", "self_attn")
|
||||
elif "attn2" in new_key:
|
||||
new_key = new_key.replace("attn2", "cross_attn")
|
||||
|
||||
# Attention components conversion
|
||||
if "to_out.0" in new_key:
|
||||
new_key = new_key.replace("to_out.0", "o")
|
||||
elif "to_q" in new_key:
|
||||
new_key = new_key.replace("to_q", "q")
|
||||
elif "to_k" in new_key:
|
||||
new_key = new_key.replace("to_k", "k")
|
||||
elif "to_v" in new_key:
|
||||
new_key = new_key.replace("to_v", "v")
|
||||
|
||||
# FFN conversion
|
||||
if "ffn.net.0.proj" in new_key:
|
||||
new_key = new_key.replace("ffn.net.0.proj", "ffn.0")
|
||||
elif "ffn.net.2" in new_key:
|
||||
new_key = new_key.replace("ffn.net.2", "ffn.2")
|
||||
|
||||
new_state_dict[new_key] = state_dict[key]
|
||||
return new_state_dict
|
||||
@@ -15,6 +15,7 @@ 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
|
||||
from optimum.quanto import QBytesTensor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from toolkit.lycoris_special import LycorisSpecialNetwork, LoConSpecialModule
|
||||
@@ -27,7 +28,8 @@ Module = Union['LoConSpecialModule', 'LoRAModule', 'DoRAModule']
|
||||
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear'
|
||||
'LoRACompatibleLinear',
|
||||
'QLinear'
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
@@ -108,11 +110,16 @@ class ExtractableModuleMixin:
|
||||
if extract_mode == "existing":
|
||||
extract_mode = 'fixed'
|
||||
extract_mode_param = self.lora_dim
|
||||
|
||||
if isinstance(weight_to_extract, QBytesTensor):
|
||||
weight_to_extract = weight_to_extract.dequantize()
|
||||
|
||||
weight_to_extract = weight_to_extract.clone().detach().float()
|
||||
|
||||
if self.org_module[0].__class__.__name__ in CONV_MODULES:
|
||||
# do conv extraction
|
||||
down_weight, up_weight, new_dim, diff = extract_conv(
|
||||
weight=weight_to_extract.clone().detach().float(),
|
||||
weight=weight_to_extract,
|
||||
mode=extract_mode,
|
||||
mode_param=extract_mode_param,
|
||||
device=device
|
||||
@@ -121,7 +128,7 @@ class ExtractableModuleMixin:
|
||||
elif self.org_module[0].__class__.__name__ in LINEAR_MODULES:
|
||||
# do linear extraction
|
||||
down_weight, up_weight, new_dim, diff = extract_linear(
|
||||
weight=weight_to_extract.clone().detach().float(),
|
||||
weight=weight_to_extract,
|
||||
mode=extract_mode,
|
||||
mode_param=extract_mode_param,
|
||||
device=device,
|
||||
@@ -175,6 +182,7 @@ class ToolkitModuleMixin:
|
||||
lx = self.lora_down(x)
|
||||
except RuntimeError as e:
|
||||
print(f"Error in {self.__class__.__name__} lora_down")
|
||||
print(e)
|
||||
|
||||
if isinstance(self.dropout, nn.Dropout) or isinstance(self.dropout, nn.Identity):
|
||||
lx = self.dropout(lx)
|
||||
@@ -209,6 +217,11 @@ class ToolkitModuleMixin:
|
||||
network: Network = self.network_ref()
|
||||
if not network.is_active:
|
||||
return self.org_forward(x, *args, **kwargs)
|
||||
|
||||
orig_dtype = x.dtype
|
||||
|
||||
if x.dtype != self.lora_down.weight.dtype:
|
||||
x = x.to(self.lora_down.weight.dtype)
|
||||
|
||||
if network.lorm_train_mode == 'local':
|
||||
# we are going to predict input with both and do a loss on them
|
||||
@@ -229,7 +242,9 @@ class ToolkitModuleMixin:
|
||||
return target_pred
|
||||
|
||||
else:
|
||||
return self.lora_up(self.lora_down(x))
|
||||
x = self.lora_up(self.lora_down(x))
|
||||
if x.dtype != orig_dtype:
|
||||
x = x.to(orig_dtype)
|
||||
|
||||
def forward(self: Module, x, *args, **kwargs):
|
||||
skip = False
|
||||
@@ -257,6 +272,9 @@ class ToolkitModuleMixin:
|
||||
# if self.__class__.__name__ == "DoRAModule":
|
||||
# # return dora forward
|
||||
# return self.dora_forward(x, *args, **kwargs)
|
||||
|
||||
if self.__class__.__name__ == "LokrModule":
|
||||
return self._call_forward(x)
|
||||
|
||||
org_forwarded = self.org_forward(x, *args, **kwargs)
|
||||
|
||||
@@ -473,13 +491,8 @@ class ToolkitNetworkMixin:
|
||||
keymap = new_keymap
|
||||
|
||||
return keymap
|
||||
|
||||
def save_weights(
|
||||
self: Network,
|
||||
file, dtype=torch.float16,
|
||||
metadata=None,
|
||||
extra_state_dict: Optional[OrderedDict] = None
|
||||
):
|
||||
|
||||
def get_state_dict(self: Network, extra_state_dict=None, dtype=torch.float16):
|
||||
keymap = self.get_keymap()
|
||||
|
||||
save_keymap = {}
|
||||
@@ -488,9 +501,6 @@ class ToolkitNetworkMixin:
|
||||
# invert them
|
||||
save_keymap[diffusers_key] = ldm_key
|
||||
|
||||
if metadata is not None and len(metadata) == 0:
|
||||
metadata = None
|
||||
|
||||
state_dict = self.state_dict()
|
||||
save_dict = OrderedDict()
|
||||
|
||||
@@ -525,10 +535,36 @@ class ToolkitNetworkMixin:
|
||||
new_save_dict[new_key] = value
|
||||
|
||||
save_dict = new_save_dict
|
||||
|
||||
|
||||
if self.network_type.lower() == "lokr":
|
||||
new_save_dict = {}
|
||||
for key, value in save_dict.items():
|
||||
# lora_transformer_transformer_blocks_7_attn_to_v.lokr_w1 to lycoris_transformer_blocks_7_attn_to_v.lokr_w1
|
||||
new_key = key
|
||||
new_key = new_key.replace('lora_transformer_', 'lycoris_')
|
||||
new_save_dict[new_key] = value
|
||||
|
||||
save_dict = new_save_dict
|
||||
|
||||
if self.base_model_ref is not None:
|
||||
save_dict = self.base_model_ref().convert_lora_weights_before_save(save_dict)
|
||||
return save_dict
|
||||
|
||||
def save_weights(
|
||||
self: Network,
|
||||
file, dtype=torch.float16,
|
||||
metadata=None,
|
||||
extra_state_dict: Optional[OrderedDict] = None
|
||||
):
|
||||
save_dict = self.get_state_dict(extra_state_dict=extra_state_dict, dtype=dtype)
|
||||
|
||||
if metadata is not None and len(metadata) == 0:
|
||||
metadata = None
|
||||
|
||||
if metadata is None:
|
||||
metadata = OrderedDict()
|
||||
metadata = add_model_hash_to_meta(state_dict, metadata)
|
||||
metadata = add_model_hash_to_meta(save_dict, metadata)
|
||||
if os.path.splitext(file)[1] == ".safetensors":
|
||||
from safetensors.torch import save_file
|
||||
save_file(save_dict, file, metadata)
|
||||
@@ -550,6 +586,9 @@ class ToolkitNetworkMixin:
|
||||
else:
|
||||
# probably a state dict
|
||||
weights_sd = file
|
||||
|
||||
if self.base_model_ref is not None:
|
||||
weights_sd = self.base_model_ref().convert_lora_weights_before_load(weights_sd)
|
||||
|
||||
load_sd = OrderedDict()
|
||||
for key, value in weights_sd.items():
|
||||
@@ -570,6 +609,10 @@ class ToolkitNetworkMixin:
|
||||
load_key = load_key.replace('.', '$$')
|
||||
load_key = load_key.replace('$$lora_down$$', '.lora_down.')
|
||||
load_key = load_key.replace('$$lora_up$$', '.lora_up.')
|
||||
|
||||
if self.network_type.lower() == "lokr":
|
||||
# lora_transformer_transformer_blocks_7_attn_to_v.lokr_w1 to lycoris_transformer_blocks_7_attn_to_v.lokr_w1
|
||||
load_key = load_key.replace('lycoris_', 'lora_transformer_')
|
||||
|
||||
load_sd[load_key] = value
|
||||
|
||||
@@ -601,9 +644,22 @@ class ToolkitNetworkMixin:
|
||||
# without having to set it in every single module every time it changes
|
||||
multiplier = self._multiplier
|
||||
# get first module
|
||||
first_module = self.get_all_modules()[0]
|
||||
device = first_module.lora_down.weight.device
|
||||
dtype = first_module.lora_down.weight.dtype
|
||||
try:
|
||||
first_module = self.get_all_modules()[0]
|
||||
except IndexError:
|
||||
raise ValueError("There are not any lora modules in this network. Check your config and try again")
|
||||
|
||||
if hasattr(first_module, 'lora_down'):
|
||||
device = first_module.lora_down.weight.device
|
||||
dtype = first_module.lora_down.weight.dtype
|
||||
elif hasattr(first_module, 'lokr_w1'):
|
||||
device = first_module.lokr_w1.device
|
||||
dtype = first_module.lokr_w1.dtype
|
||||
elif hasattr(first_module, 'lokr_w1_a'):
|
||||
device = first_module.lokr_w1_a.device
|
||||
dtype = first_module.lokr_w1_a.dtype
|
||||
else:
|
||||
raise ValueError("Unknown module type")
|
||||
with torch.no_grad():
|
||||
tensor_multiplier = None
|
||||
if isinstance(multiplier, int) or isinstance(multiplier, float):
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import torch
|
||||
from transformers import Adafactor, AdamW
|
||||
|
||||
|
||||
def get_optimizer(
|
||||
@@ -28,6 +27,18 @@ def get_optimizer(
|
||||
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("prodigy8bit"):
|
||||
from toolkit.optimizers.prodigy_8bit import Prodigy8bit
|
||||
print("Using Prodigy optimizer")
|
||||
use_lr = learning_rate
|
||||
if use_lr < 0.1:
|
||||
# dadaptation uses different lr that is values of 0.1 to 1.0. default to 1.0
|
||||
use_lr = 1.0
|
||||
|
||||
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 = Prodigy8bit(params, lr=use_lr, eps=1e-6, **optimizer_params)
|
||||
elif lower_type.startswith("prodigy"):
|
||||
from prodigyopt import Prodigy
|
||||
|
||||
@@ -41,11 +52,21 @@ def get_optimizer(
|
||||
# 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, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "adam8":
|
||||
from toolkit.optimizers.adam8bit import Adam8bit
|
||||
|
||||
optimizer = Adam8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "adamw8":
|
||||
from toolkit.optimizers.adam8bit import Adam8bit
|
||||
|
||||
optimizer = Adam8bit(params, lr=learning_rate, eps=1e-6, decouple=True, **optimizer_params)
|
||||
elif lower_type.endswith("8bit"):
|
||||
import bitsandbytes
|
||||
|
||||
if lower_type == "adam8bit":
|
||||
return bitsandbytes.optim.Adam8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
if lower_type == "ademamix8bit":
|
||||
return bitsandbytes.optim.AdEMAMix8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "adamw8bit":
|
||||
return bitsandbytes.optim.AdamW8bit(params, lr=learning_rate, eps=1e-6, **optimizer_params)
|
||||
elif lower_type == "lion8bit":
|
||||
@@ -63,9 +84,9 @@ 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), eps=1e-6, **optimizer_params)
|
||||
optimizer = torch.optim.Adagrad(params, lr=float(learning_rate), **optimizer_params)
|
||||
elif lower_type == 'adafactor':
|
||||
# hack in stochastic rounding
|
||||
from toolkit.optimizers.adafactor import Adafactor
|
||||
if 'relative_step' not in optimizer_params:
|
||||
optimizer_params['relative_step'] = False
|
||||
if 'scale_parameter' not in optimizer_params:
|
||||
@@ -73,8 +94,9 @@ def get_optimizer(
|
||||
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)
|
||||
elif lower_type == 'automagic':
|
||||
from toolkit.optimizers.automagic import Automagic
|
||||
optimizer = Automagic(params, lr=float(learning_rate), **optimizer_params)
|
||||
else:
|
||||
raise ValueError(f'Unknown optimizer type {optimizer_type}')
|
||||
return optimizer
|
||||
|
||||
361
toolkit/optimizers/adafactor.py
Normal file
361
toolkit/optimizers/adafactor.py
Normal file
@@ -0,0 +1,361 @@
|
||||
import math
|
||||
from typing import List
|
||||
import torch
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, stochastic_grad_accummulation
|
||||
from optimum.quanto import QBytesTensor
|
||||
import random
|
||||
|
||||
|
||||
class Adafactor(torch.optim.Optimizer):
|
||||
"""
|
||||
Adafactor implementation with stochastic rounding accumulation and stochastic rounding on apply.
|
||||
Modified from transformers Adafactor implementation to support stochastic rounding accumulation and apply.
|
||||
|
||||
AdaFactor pytorch implementation can be used as a drop in replacement for Adam original fairseq code:
|
||||
https://github.com/pytorch/fairseq/blob/master/fairseq/optim/adafactor.py
|
||||
|
||||
Paper: *Adafactor: Adaptive Learning Rates with Sublinear Memory Cost* https://arxiv.org/abs/1804.04235 Note that
|
||||
this optimizer internally adjusts the learning rate depending on the `scale_parameter`, `relative_step` and
|
||||
`warmup_init` options. To use a manual (external) learning rate schedule you should set `scale_parameter=False` and
|
||||
`relative_step=False`.
|
||||
|
||||
Arguments:
|
||||
params (`Iterable[nn.parameter.Parameter]`):
|
||||
Iterable of parameters to optimize or dictionaries defining parameter groups.
|
||||
lr (`float`, *optional*):
|
||||
The external learning rate.
|
||||
eps (`Tuple[float, float]`, *optional*, defaults to `(1e-30, 0.001)`):
|
||||
Regularization constants for square gradient and parameter scale respectively
|
||||
clip_threshold (`float`, *optional*, defaults to 1.0):
|
||||
Threshold of root mean square of final gradient update
|
||||
decay_rate (`float`, *optional*, defaults to -0.8):
|
||||
Coefficient used to compute running averages of square
|
||||
beta1 (`float`, *optional*):
|
||||
Coefficient used for computing running averages of gradient
|
||||
weight_decay (`float`, *optional*, defaults to 0.0):
|
||||
Weight decay (L2 penalty)
|
||||
scale_parameter (`bool`, *optional*, defaults to `True`):
|
||||
If True, learning rate is scaled by root mean square
|
||||
relative_step (`bool`, *optional*, defaults to `True`):
|
||||
If True, time-dependent learning rate is computed instead of external learning rate
|
||||
warmup_init (`bool`, *optional*, defaults to `False`):
|
||||
Time-dependent learning rate computation depends on whether warm-up initialization is being used
|
||||
|
||||
This implementation handles low-precision (FP16, bfloat) values, but we have not thoroughly tested.
|
||||
|
||||
Recommended T5 finetuning settings (https://discuss.huggingface.co/t/t5-finetuning-tips/684/3):
|
||||
|
||||
- Training without LR warmup or clip_threshold is not recommended.
|
||||
|
||||
- use scheduled LR warm-up to fixed LR
|
||||
- use clip_threshold=1.0 (https://arxiv.org/abs/1804.04235)
|
||||
- Disable relative updates
|
||||
- Use scale_parameter=False
|
||||
- Additional optimizer operations like gradient clipping should not be used alongside Adafactor
|
||||
|
||||
Example:
|
||||
|
||||
```python
|
||||
Adafactor(model.parameters(), scale_parameter=False, relative_step=False, warmup_init=False, lr=1e-3)
|
||||
```
|
||||
|
||||
Others reported the following combination to work well:
|
||||
|
||||
```python
|
||||
Adafactor(model.parameters(), scale_parameter=True, relative_step=True, warmup_init=True, lr=None)
|
||||
```
|
||||
|
||||
When using `lr=None` with [`Trainer`] you will most likely need to use [`~optimization.AdafactorSchedule`]
|
||||
scheduler as following:
|
||||
|
||||
```python
|
||||
from transformers.optimization import Adafactor, AdafactorSchedule
|
||||
|
||||
optimizer = Adafactor(model.parameters(), scale_parameter=True, relative_step=True, warmup_init=True, lr=None)
|
||||
lr_scheduler = AdafactorSchedule(optimizer)
|
||||
trainer = Trainer(..., optimizers=(optimizer, lr_scheduler))
|
||||
```
|
||||
|
||||
Usage:
|
||||
|
||||
```python
|
||||
# replace AdamW with Adafactor
|
||||
optimizer = Adafactor(
|
||||
model.parameters(),
|
||||
lr=1e-3,
|
||||
eps=(1e-30, 1e-3),
|
||||
clip_threshold=1.0,
|
||||
decay_rate=-0.8,
|
||||
beta1=None,
|
||||
weight_decay=0.0,
|
||||
relative_step=False,
|
||||
scale_parameter=False,
|
||||
warmup_init=False,
|
||||
)
|
||||
```"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr=None,
|
||||
eps=(1e-30, 1e-3),
|
||||
clip_threshold=1.0,
|
||||
decay_rate=-0.8,
|
||||
beta1=None,
|
||||
weight_decay=0.0,
|
||||
scale_parameter=True,
|
||||
relative_step=True,
|
||||
warmup_init=False,
|
||||
do_paramiter_swapping=False,
|
||||
paramiter_swapping_factor=0.1,
|
||||
stochastic_accumulation=True,
|
||||
):
|
||||
if lr is not None and relative_step:
|
||||
raise ValueError(
|
||||
"Cannot combine manual `lr` and `relative_step=True` options")
|
||||
if warmup_init and not relative_step:
|
||||
raise ValueError(
|
||||
"`warmup_init=True` requires `relative_step=True`")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"eps": eps,
|
||||
"clip_threshold": clip_threshold,
|
||||
"decay_rate": decay_rate,
|
||||
"beta1": beta1,
|
||||
"weight_decay": weight_decay,
|
||||
"scale_parameter": scale_parameter,
|
||||
"relative_step": relative_step,
|
||||
"warmup_init": warmup_init,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
|
||||
self.base_lrs: List[float] = [
|
||||
lr for group in self.param_groups
|
||||
]
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# setup stochastic grad accum hooks
|
||||
if stochastic_accumulation:
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
self.do_paramiter_swapping = do_paramiter_swapping
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
self._total_paramiter_size = 0
|
||||
# count total paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
self._total_paramiter_size += torch.numel(param)
|
||||
# pretty print total paramiters with comma seperation
|
||||
print(f"Total training paramiters: {self._total_paramiter_size:,}")
|
||||
|
||||
# needs to be enabled to count paramiters
|
||||
if self.do_paramiter_swapping:
|
||||
self.enable_paramiter_swapping(self.paramiter_swapping_factor)
|
||||
|
||||
|
||||
def enable_paramiter_swapping(self, paramiter_swapping_factor=0.1):
|
||||
self.do_paramiter_swapping = True
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
# call it an initial time
|
||||
self.swap_paramiters()
|
||||
|
||||
def swap_paramiters(self):
|
||||
all_params = []
|
||||
# deactivate all paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
param.requires_grad_(False)
|
||||
# remove any grad
|
||||
param.grad = None
|
||||
all_params.append(param)
|
||||
# shuffle all paramiters
|
||||
random.shuffle(all_params)
|
||||
|
||||
# keep activating paramiters until we are going to go over the target paramiters
|
||||
target_paramiters = int(self._total_paramiter_size * self.paramiter_swapping_factor)
|
||||
total_paramiters = 0
|
||||
for param in all_params:
|
||||
total_paramiters += torch.numel(param)
|
||||
if total_paramiters >= target_paramiters:
|
||||
break
|
||||
else:
|
||||
param.requires_grad_(True)
|
||||
|
||||
@staticmethod
|
||||
def _get_lr(param_group, param_state):
|
||||
rel_step_sz = param_group["lr"]
|
||||
if param_group["relative_step"]:
|
||||
min_step = 1e-6 * \
|
||||
param_state["step"] if param_group["warmup_init"] else 1e-2
|
||||
rel_step_sz = min(min_step, 1.0 / math.sqrt(param_state["step"]))
|
||||
param_scale = 1.0
|
||||
if param_group["scale_parameter"]:
|
||||
param_scale = max(param_group["eps"][1], param_state["RMS"])
|
||||
return param_scale * rel_step_sz
|
||||
|
||||
@staticmethod
|
||||
def _get_options(param_group, param_shape):
|
||||
factored = len(param_shape) >= 2
|
||||
use_first_moment = param_group["beta1"] is not None
|
||||
return factored, use_first_moment
|
||||
|
||||
@staticmethod
|
||||
def _rms(tensor):
|
||||
return tensor.norm(2) / (tensor.numel() ** 0.5)
|
||||
|
||||
@staticmethod
|
||||
def _approx_sq_grad(exp_avg_sq_row, exp_avg_sq_col):
|
||||
# copy from fairseq's adafactor implementation:
|
||||
# https://github.com/huggingface/transformers/blob/8395f14de6068012787d83989c3627c3df6a252b/src/transformers/optimization.py#L505
|
||||
r_factor = (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-
|
||||
1, keepdim=True)).rsqrt_().unsqueeze(-1)
|
||||
c_factor = exp_avg_sq_col.unsqueeze(-2).rsqrt()
|
||||
return torch.mul(r_factor, c_factor)
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
# adafactor manages its own lr
|
||||
def get_learning_rates(self):
|
||||
lrs = [
|
||||
self._get_lr(group, self.state[group["params"][0]])
|
||||
for group in self.param_groups
|
||||
if group["params"][0].grad is not None
|
||||
]
|
||||
if len(lrs) == 0:
|
||||
lrs = self.base_lrs # if called before stepping
|
||||
return lrs
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""
|
||||
Performs a single optimization step
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
self.step_hook()
|
||||
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 or not p.requires_grad:
|
||||
continue
|
||||
|
||||
grad = p.grad
|
||||
if grad.dtype != torch.float32:
|
||||
grad = grad.to(torch.float32)
|
||||
if grad.is_sparse:
|
||||
raise RuntimeError(
|
||||
"Adafactor does not support sparse gradients.")
|
||||
|
||||
# if p has atts _scale then it is quantized. We need to divide the grad by the scale
|
||||
# if hasattr(p, "_scale"):
|
||||
# grad = grad / p._scale
|
||||
|
||||
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 isinstance(p_data_fp32, QBytesTensor):
|
||||
p_data_fp32 = p_data_fp32.dequantize()
|
||||
if p.dtype != torch.float32:
|
||||
p_data_fp32 = p_data_fp32.clone().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"]
|
||||
if isinstance(eps, tuple) or isinstance(eps, list):
|
||||
eps = eps[0]
|
||||
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.float32:
|
||||
# apply stochastic rounding
|
||||
copy_stochastic(p, p_data_fp32)
|
||||
|
||||
return loss
|
||||
162
toolkit/optimizers/adam8bit.py
Normal file
162
toolkit/optimizers/adam8bit.py
Normal file
@@ -0,0 +1,162 @@
|
||||
import math
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, Auto8bitTensor, stochastic_grad_accummulation
|
||||
|
||||
class Adam8bit(Optimizer):
|
||||
"""
|
||||
Implements Adam optimizer with 8-bit state storage and stochastic rounding.
|
||||
|
||||
Arguments:
|
||||
params (iterable): Iterable of parameters to optimize or dicts defining parameter groups
|
||||
lr (float): Learning rate (default: 1e-3)
|
||||
betas (tuple): Coefficients for computing running averages of gradient and its square (default: (0.9, 0.999))
|
||||
eps (float): Term added to denominator to improve numerical stability (default: 1e-8)
|
||||
weight_decay (float): Weight decay coefficient (default: 0)
|
||||
decouple (bool): Use AdamW style decoupled weight decay (default: True)
|
||||
"""
|
||||
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
|
||||
weight_decay=0, decouple=True):
|
||||
if not 0.0 <= lr:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if not 0.0 <= eps:
|
||||
raise ValueError(f"Invalid epsilon value: {eps}")
|
||||
if not 0.0 <= betas[0] < 1.0:
|
||||
raise ValueError(f"Invalid beta parameter at index 0: {betas[0]}")
|
||||
if not 0.0 <= betas[1] < 1.0:
|
||||
raise ValueError(f"Invalid beta parameter at index 1: {betas[1]}")
|
||||
|
||||
defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay,
|
||||
decouple=decouple)
|
||||
super(Adam8bit, self).__init__(params, defaults)
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# Setup stochastic grad accumulation hooks
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_memory_efficient_fp16(self):
|
||||
return False
|
||||
|
||||
@property
|
||||
def supports_flat_params(self):
|
||||
return True
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# Copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""Performs a single optimization step.
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model and returns the loss.
|
||||
"""
|
||||
# Call pre step
|
||||
self.step_hook()
|
||||
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
beta1, beta2 = group['betas']
|
||||
eps = group['eps']
|
||||
lr = group['lr']
|
||||
decay = group['weight_decay']
|
||||
decouple = group['decouple']
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
# Apply weight decay (coupled variant)
|
||||
if decay != 0 and not decouple:
|
||||
grad.add_(p_fp32.data, alpha=decay)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
# State initialization
|
||||
if len(state) == 0:
|
||||
state['step'] = 0
|
||||
# Exponential moving average of gradient values
|
||||
state['exp_avg'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
# Exponential moving average of squared gradient values
|
||||
state['exp_avg_sq'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
state['step'] += 1
|
||||
bias_correction1 = 1 - beta1 ** state['step']
|
||||
bias_correction2 = 1 - beta2 ** state['step']
|
||||
|
||||
# Adam EMA updates
|
||||
exp_avg.mul_(beta1).add_(grad, alpha=1-beta1)
|
||||
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1-beta2)
|
||||
|
||||
# Apply weight decay (decoupled variant)
|
||||
if decay != 0 and decouple:
|
||||
p_fp32.data.mul_(1 - lr * decay)
|
||||
|
||||
# Bias correction
|
||||
step_size = lr / bias_correction1
|
||||
denom = (exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps)
|
||||
|
||||
# Take step
|
||||
p_fp32.data.addcdiv_(exp_avg, denom, value=-step_size)
|
||||
|
||||
# Update state with stochastic rounding
|
||||
state['exp_avg'] = Auto8bitTensor(exp_avg)
|
||||
state['exp_avg_sq'] = Auto8bitTensor(exp_avg_sq)
|
||||
|
||||
# Apply stochastic rounding to parameters
|
||||
copy_stochastic(p.data, p_fp32.data)
|
||||
|
||||
return loss
|
||||
|
||||
def state_dict(self):
|
||||
"""Returns the state of the optimizer as a dict."""
|
||||
state_dict = super().state_dict()
|
||||
|
||||
# Convert Auto8bitTensor objects to regular state dicts
|
||||
for param_id, param_state in state_dict['state'].items():
|
||||
for key, value in param_state.items():
|
||||
if isinstance(value, Auto8bitTensor):
|
||||
param_state[key] = {
|
||||
'_type': 'Auto8bitTensor',
|
||||
'state': value.state_dict()
|
||||
}
|
||||
|
||||
return state_dict
|
||||
|
||||
def load_state_dict(self, state_dict):
|
||||
"""Loads the optimizer state."""
|
||||
# First, load the basic state
|
||||
super().load_state_dict(state_dict)
|
||||
|
||||
# Then convert any Auto8bitTensor states back to objects
|
||||
for param_id, param_state in self.state.items():
|
||||
for key, value in param_state.items():
|
||||
if isinstance(value, dict) and value.get('_type') == 'Auto8bitTensor':
|
||||
param_state[key] = Auto8bitTensor(value['state'])
|
||||
|
||||
337
toolkit/optimizers/automagic.py
Normal file
337
toolkit/optimizers/automagic.py
Normal file
@@ -0,0 +1,337 @@
|
||||
from collections import OrderedDict
|
||||
import math
|
||||
from typing import List
|
||||
import torch
|
||||
from toolkit.optimizers.optimizer_utils import Auto8bitTensor, copy_stochastic, stochastic_grad_accummulation
|
||||
from optimum.quanto import QBytesTensor
|
||||
import random
|
||||
|
||||
|
||||
class Automagic(torch.optim.Optimizer):
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr=None,
|
||||
min_lr=1e-7,
|
||||
max_lr=1e-3,
|
||||
lr_pump_scale=1.1,
|
||||
lr_dump_scale=0.85,
|
||||
eps=(1e-30, 1e-3),
|
||||
clip_threshold=1.0,
|
||||
decay_rate=-0.8,
|
||||
weight_decay=0.0,
|
||||
do_paramiter_swapping=False,
|
||||
paramiter_swapping_factor=0.1,
|
||||
):
|
||||
self.lr = lr
|
||||
self.min_lr = min_lr
|
||||
self.max_lr = max_lr
|
||||
self.lr_pump_scale = lr_pump_scale
|
||||
self.lr_dump_scale = lr_dump_scale
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"eps": eps,
|
||||
"clip_threshold": clip_threshold,
|
||||
"decay_rate": decay_rate,
|
||||
"weight_decay": weight_decay,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
|
||||
self.base_lrs: List[float] = [
|
||||
lr for group in self.param_groups
|
||||
]
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# setup stochastic grad accum hooks
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
self.do_paramiter_swapping = do_paramiter_swapping
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
self._total_paramiter_size = 0
|
||||
# count total paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
self._total_paramiter_size += torch.numel(param)
|
||||
# pretty print total paramiters with comma seperation
|
||||
print(f"Total training paramiters: {self._total_paramiter_size:,}")
|
||||
|
||||
# needs to be enabled to count paramiters
|
||||
if self.do_paramiter_swapping:
|
||||
self.enable_paramiter_swapping(self.paramiter_swapping_factor)
|
||||
|
||||
def enable_paramiter_swapping(self, paramiter_swapping_factor=0.1):
|
||||
self.do_paramiter_swapping = True
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
# call it an initial time
|
||||
self.swap_paramiters()
|
||||
|
||||
def swap_paramiters(self):
|
||||
all_params = []
|
||||
# deactivate all paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
param.requires_grad_(False)
|
||||
# remove any grad
|
||||
param.grad = None
|
||||
all_params.append(param)
|
||||
# shuffle all paramiters
|
||||
random.shuffle(all_params)
|
||||
|
||||
# keep activating paramiters until we are going to go over the target paramiters
|
||||
target_paramiters = int(
|
||||
self._total_paramiter_size * self.paramiter_swapping_factor)
|
||||
total_paramiters = 0
|
||||
for param in all_params:
|
||||
total_paramiters += torch.numel(param)
|
||||
if total_paramiters >= target_paramiters:
|
||||
break
|
||||
else:
|
||||
param.requires_grad_(True)
|
||||
|
||||
@staticmethod
|
||||
def _get_lr(param_group, param_state):
|
||||
if 'avg_lr' in param_state:
|
||||
lr = param_state["avg_lr"]
|
||||
else:
|
||||
lr = 0.0
|
||||
return lr
|
||||
|
||||
def _get_group_lr(self, group):
|
||||
group_lrs = []
|
||||
for p in group["params"]:
|
||||
group_lrs.append(self._get_lr(group, self.state[p]))
|
||||
# return avg
|
||||
if len(group_lrs) == 0:
|
||||
return self.lr
|
||||
return sum(group_lrs) / len(group_lrs)
|
||||
|
||||
@staticmethod
|
||||
def _rms(tensor):
|
||||
return tensor.norm(2) / (tensor.numel() ** 0.5)
|
||||
|
||||
@staticmethod
|
||||
def _approx_sq_grad(exp_avg_sq_row, exp_avg_sq_col):
|
||||
# copy from fairseq's adafactor implementation:
|
||||
# https://github.com/huggingface/transformers/blob/8395f14de6068012787d83989c3627c3df6a252b/src/transformers/optimization.py#L505
|
||||
r_factor = (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-
|
||||
1, keepdim=True)).rsqrt_().unsqueeze(-1)
|
||||
c_factor = exp_avg_sq_col.unsqueeze(-2).rsqrt()
|
||||
return torch.mul(r_factor, c_factor)
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
# adafactor manages its own lr
|
||||
def get_learning_rates(self):
|
||||
|
||||
lrs = [
|
||||
self._get_group_lr(group)
|
||||
for group in self.param_groups
|
||||
]
|
||||
if len(lrs) == 0:
|
||||
lrs = self.base_lrs # if called before stepping
|
||||
return lrs
|
||||
|
||||
def get_avg_learning_rate(self):
|
||||
lrs = self.get_learning_rates()
|
||||
return sum(lrs) / len(lrs)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""
|
||||
Performs a single optimization step
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
self.step_hook()
|
||||
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 or not p.requires_grad:
|
||||
continue
|
||||
|
||||
grad = p.grad
|
||||
if grad.dtype != torch.float32:
|
||||
grad = grad.to(torch.float32)
|
||||
if grad.is_sparse:
|
||||
raise RuntimeError(
|
||||
"Automagic does not support sparse gradients.")
|
||||
|
||||
state = self.state[p]
|
||||
grad_shape = grad.shape
|
||||
|
||||
factored = len(grad_shape) >= 2
|
||||
# State Initialization
|
||||
if len(state) == 0:
|
||||
self.initialize_state(p)
|
||||
else:
|
||||
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 isinstance(p_data_fp32, QBytesTensor):
|
||||
p_data_fp32 = p_data_fp32.dequantize()
|
||||
if p.dtype != torch.float32:
|
||||
p_data_fp32 = p_data_fp32.clone().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"]
|
||||
if isinstance(eps, tuple) or isinstance(eps, list):
|
||||
eps = eps[0]
|
||||
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))
|
||||
|
||||
# calculate new lr mask. if the updated param is going in same direction, increase lr, else decrease
|
||||
# update the lr mask. self.lr_momentum is < 1.0. If a paramiter is positive and increasing (or negative and decreasing), increase lr,
|
||||
# for that single paramiter. If a paramiter is negative and increasing or positive and decreasing, decrease lr for that single paramiter.
|
||||
# to decrease lr, multiple by self.lr_momentum, to increase lr, divide by self.lr_momentum.
|
||||
|
||||
# not doing it this way anymore
|
||||
# update.mul_(lr)
|
||||
|
||||
# Get signs of current last update and updates
|
||||
last_polarity = state['last_polarity']
|
||||
current_polarity = (update > 0).to(torch.bool)
|
||||
sign_agreement = torch.where(
|
||||
last_polarity == current_polarity, 1, -1)
|
||||
state['last_polarity'] = current_polarity
|
||||
|
||||
lr_mask = state['lr_mask'].to(torch.float32)
|
||||
|
||||
# Update learning rate mask based on sign agreement
|
||||
new_lr = torch.where(
|
||||
sign_agreement > 0,
|
||||
lr_mask * self.lr_pump_scale, # Increase lr
|
||||
lr_mask * self.lr_dump_scale # Decrease lr
|
||||
)
|
||||
|
||||
# Clip learning rates to bounds
|
||||
new_lr = torch.clamp(
|
||||
new_lr,
|
||||
min=self.min_lr,
|
||||
max=self.max_lr
|
||||
)
|
||||
|
||||
# Apply the learning rate mask to the update
|
||||
update.mul_(new_lr)
|
||||
|
||||
state['lr_mask'] = Auto8bitTensor(new_lr)
|
||||
state['avg_lr'] = torch.mean(new_lr)
|
||||
|
||||
if group["weight_decay"] != 0:
|
||||
p_data_fp32.add_(
|
||||
p_data_fp32, alpha=(-group["weight_decay"] * new_lr))
|
||||
|
||||
p_data_fp32.add_(-update)
|
||||
|
||||
if p.dtype != torch.float32:
|
||||
# apply stochastic rounding
|
||||
copy_stochastic(p, p_data_fp32)
|
||||
|
||||
return loss
|
||||
|
||||
def initialize_state(self, p):
|
||||
state = self.state[p]
|
||||
state["step"] = 0
|
||||
|
||||
# store the lr mask
|
||||
if 'lr_mask' not in state:
|
||||
state['lr_mask'] = Auto8bitTensor(torch.ones(
|
||||
p.shape).to(p.device, dtype=torch.float32) * self.lr
|
||||
)
|
||||
state['avg_lr'] = torch.mean(
|
||||
state['lr_mask'].to(torch.float32))
|
||||
if 'last_polarity' not in state:
|
||||
state['last_polarity'] = torch.zeros(
|
||||
p.shape, dtype=torch.bool, device=p.device)
|
||||
|
||||
factored = len(p.shape) >= 2
|
||||
if factored:
|
||||
state["exp_avg_sq_row"] = torch.zeros(
|
||||
p.shape[:-1]).to(p)
|
||||
state["exp_avg_sq_col"] = torch.zeros(
|
||||
p.shape[:-2] + p.shape[-1:]).to(p)
|
||||
else:
|
||||
state["exp_avg_sq"] = torch.zeros_like(p)
|
||||
|
||||
state["RMS"] = 0
|
||||
|
||||
# override the state_dict to save the lr_mask
|
||||
def state_dict(self, *args, **kwargs):
|
||||
orig_state_dict = super().state_dict(*args, **kwargs)
|
||||
# convert the state to quantized tensor to scale and quantized
|
||||
new_sace_state = {}
|
||||
for p, state in orig_state_dict['state'].items():
|
||||
save_state = {k: v for k, v in state.items() if k != 'lr_mask'}
|
||||
save_state['lr_mask'] = state['lr_mask'].state_dict()
|
||||
new_sace_state[p] = save_state
|
||||
|
||||
orig_state_dict['state'] = new_sace_state
|
||||
|
||||
return orig_state_dict
|
||||
|
||||
def load_state_dict(self, state_dict, strict=True):
|
||||
# load the lr_mask from the state_dict
|
||||
# dont load state dict for now. Has a bug. Need to fix it.
|
||||
return
|
||||
idx = 0
|
||||
for group in self.param_groups:
|
||||
for p in group['params']:
|
||||
self.initialize_state(p)
|
||||
state = self.state[p]
|
||||
m = state_dict['state'][idx]['lr_mask']
|
||||
sd_mask = m['quantized'].to(m['orig_dtype']) * m['scale']
|
||||
state['lr_mask'] = Auto8bitTensor(sd_mask)
|
||||
del state_dict['state'][idx]['lr_mask']
|
||||
idx += 1
|
||||
super().load_state_dict(state_dict)
|
||||
256
toolkit/optimizers/optimizer_utils.py
Normal file
256
toolkit/optimizers/optimizer_utils.py
Normal file
@@ -0,0 +1,256 @@
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from typing import Optional
|
||||
from optimum.quanto import QBytesTensor
|
||||
|
||||
|
||||
def compute_scale_for_dtype(tensor, dtype):
|
||||
"""
|
||||
Compute appropriate scale for the given tensor and target dtype.
|
||||
|
||||
Args:
|
||||
tensor: Input tensor to be quantized
|
||||
dtype: Target dtype for quantization
|
||||
Returns:
|
||||
Appropriate scale factor for the quantization
|
||||
"""
|
||||
if dtype == torch.int8:
|
||||
abs_max = torch.max(torch.abs(tensor))
|
||||
return abs_max / 127.0 if abs_max > 0 else 1.0
|
||||
elif dtype == torch.uint8:
|
||||
max_val = torch.max(tensor)
|
||||
min_val = torch.min(tensor)
|
||||
range_val = max_val - min_val
|
||||
return range_val / 255.0 if range_val > 0 else 1.0
|
||||
elif dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
# For float8, we typically want to preserve the magnitude of the values
|
||||
# while fitting within the representable range of the format
|
||||
abs_max = torch.max(torch.abs(tensor))
|
||||
if dtype == torch.float8_e4m3fn:
|
||||
# e4m3fn has range [-448, 448] with no infinities
|
||||
max_representable = 448.0
|
||||
else: # torch.float8_e5m2
|
||||
# e5m2 has range [-57344, 57344] with infinities
|
||||
max_representable = 57344.0
|
||||
|
||||
return abs_max / max_representable if abs_max > 0 else 1.0
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype for quantization: {dtype}")
|
||||
|
||||
def quantize_tensor(tensor, dtype):
|
||||
"""
|
||||
Quantize a floating-point tensor to the target dtype with appropriate scaling.
|
||||
|
||||
Args:
|
||||
tensor: Input tensor (float)
|
||||
dtype: Target dtype for quantization
|
||||
Returns:
|
||||
quantized_data: Quantized tensor
|
||||
scale: Scale factor used
|
||||
"""
|
||||
scale = compute_scale_for_dtype(tensor, dtype)
|
||||
|
||||
if dtype == torch.int8:
|
||||
quantized_data = torch.clamp(torch.round(tensor / scale), -128, 127).to(dtype)
|
||||
elif dtype == torch.uint8:
|
||||
quantized_data = torch.clamp(torch.round(tensor / scale), 0, 255).to(dtype)
|
||||
elif dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
# For float8, we scale and then cast directly to the target type
|
||||
# The casting operation will handle the appropriate rounding
|
||||
scaled_tensor = tensor / scale
|
||||
quantized_data = scaled_tensor.to(dtype)
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype for quantization: {dtype}")
|
||||
|
||||
return quantized_data, scale
|
||||
|
||||
|
||||
def update_parameter(target, result_float):
|
||||
"""
|
||||
Updates a parameter tensor, handling both regular torch.Tensor and QBytesTensor cases
|
||||
with proper rescaling for quantized tensors.
|
||||
|
||||
Args:
|
||||
target: The parameter to update (either torch.Tensor or QBytesTensor)
|
||||
result_float: The new values to assign (torch.Tensor)
|
||||
"""
|
||||
if isinstance(target, QBytesTensor):
|
||||
# Get the target dtype from the existing quantized tensor
|
||||
target_dtype = target._data.dtype
|
||||
|
||||
# Handle device placement
|
||||
device = target._data.device
|
||||
result_float = result_float.to(device)
|
||||
|
||||
# Compute new quantized values and scale
|
||||
quantized_data, new_scale = quantize_tensor(result_float, target_dtype)
|
||||
|
||||
# Update the internal tensors with newly computed values
|
||||
target._data.copy_(quantized_data)
|
||||
target._scale.copy_(new_scale)
|
||||
else:
|
||||
# Regular tensor update
|
||||
target.copy_(result_float)
|
||||
|
||||
|
||||
def get_format_params(dtype: torch.dtype) -> tuple[int, int]:
|
||||
"""
|
||||
Returns (mantissa_bits, total_bits) for each format.
|
||||
mantissa_bits excludes the implicit leading 1.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
return 23, 32
|
||||
elif dtype == torch.bfloat16:
|
||||
return 7, 16
|
||||
elif dtype == torch.float16:
|
||||
return 10, 16
|
||||
elif dtype == torch.float8_e4m3fn:
|
||||
return 3, 8
|
||||
elif dtype == torch.float8_e5m2:
|
||||
return 2, 8
|
||||
elif dtype == torch.int8:
|
||||
return 0, 8 # Int8 doesn't have mantissa bits
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {dtype}")
|
||||
|
||||
|
||||
def copy_stochastic(
|
||||
target: torch.Tensor,
|
||||
source: torch.Tensor,
|
||||
eps: Optional[float] = None
|
||||
) -> None:
|
||||
"""
|
||||
Performs stochastic rounding from source tensor to target tensor.
|
||||
|
||||
Args:
|
||||
target: Destination tensor (determines the target format)
|
||||
source: Source tensor (typically float32)
|
||||
eps: Optional minimum value for stochastic rounding (for numerical stability)
|
||||
"""
|
||||
with torch.no_grad():
|
||||
# If target is float32, just copy directly
|
||||
if target.dtype == torch.float32:
|
||||
target.copy_(source)
|
||||
return
|
||||
|
||||
# Special handling for int8
|
||||
if target.dtype == torch.int8:
|
||||
# Scale the source values to utilize the full int8 range
|
||||
scaled = source * 127.0 # Scale to [-127, 127]
|
||||
|
||||
# Add random noise for stochastic rounding
|
||||
noise = torch.rand_like(scaled) - 0.5
|
||||
rounded = torch.round(scaled + noise)
|
||||
|
||||
# Clamp to int8 range
|
||||
clamped = torch.clamp(rounded, -127, 127)
|
||||
target.copy_(clamped.to(torch.int8))
|
||||
return
|
||||
|
||||
mantissa_bits, _ = get_format_params(target.dtype)
|
||||
|
||||
# Convert source to int32 view
|
||||
source_int = source.view(dtype=torch.int32)
|
||||
|
||||
# Calculate number of bits to round
|
||||
bits_to_round = 23 - mantissa_bits # 23 is float32 mantissa bits
|
||||
|
||||
# Create random integers for stochastic rounding
|
||||
rand = torch.randint_like(
|
||||
source,
|
||||
dtype=torch.int32,
|
||||
low=0,
|
||||
high=(1 << bits_to_round),
|
||||
)
|
||||
|
||||
# Add random values to the bits that will be rounded off
|
||||
result = source_int.clone()
|
||||
result.add_(rand)
|
||||
|
||||
# Mask to keep only the bits we want
|
||||
# Create mask with 1s in positions we want to keep
|
||||
mask = (-1) << bits_to_round
|
||||
result.bitwise_and_(mask)
|
||||
|
||||
# Handle minimum value threshold if specified
|
||||
if eps is not None:
|
||||
eps_int = torch.tensor(
|
||||
eps, dtype=torch.float32).view(dtype=torch.int32)
|
||||
zero_mask = (result.abs() < eps_int)
|
||||
result[zero_mask] = torch.sign(source_int[zero_mask]) * eps_int
|
||||
|
||||
# Convert back to float32 view
|
||||
result_float = result.view(dtype=torch.float32)
|
||||
|
||||
# Special handling for float8 formats
|
||||
if target.dtype == torch.float8_e4m3fn:
|
||||
result_float.clamp_(-448.0, 448.0)
|
||||
elif target.dtype == torch.float8_e5m2:
|
||||
result_float.clamp_(-57344.0, 57344.0)
|
||||
|
||||
# Copy the result to the target tensor
|
||||
update_parameter(target, result_float)
|
||||
# target.copy_(result_float)
|
||||
del result, rand, source_int
|
||||
|
||||
|
||||
class Auto8bitTensor:
|
||||
def __init__(self, data: Tensor, *args, **kwargs):
|
||||
if isinstance(data, dict): # Add constructor from state dict
|
||||
self._load_from_state_dict(data)
|
||||
else:
|
||||
abs_max = data.abs().max().item()
|
||||
scale = abs_max / 127.0 if abs_max > 0 else 1.0
|
||||
|
||||
self.quantized = (data / scale).round().clamp(-127, 127).to(torch.int8)
|
||||
self.scale = scale
|
||||
self.orig_dtype = data.dtype
|
||||
|
||||
def dequantize(self) -> Tensor:
|
||||
return self.quantized.to(dtype=torch.float32) * self.scale
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
# Handle the dtype argument whether it's positional or keyword
|
||||
dtype = None
|
||||
if args and isinstance(args[0], torch.dtype):
|
||||
dtype = args[0]
|
||||
args = args[1:]
|
||||
elif 'dtype' in kwargs:
|
||||
dtype = kwargs['dtype']
|
||||
del kwargs['dtype']
|
||||
|
||||
if dtype is not None:
|
||||
# First dequantize then convert to requested dtype
|
||||
return self.dequantize().to(dtype=dtype, *args, **kwargs)
|
||||
|
||||
# If no dtype specified, just pass through to parent
|
||||
return self.dequantize().to(*args, **kwargs)
|
||||
|
||||
def state_dict(self):
|
||||
"""Returns a dictionary containing the current state of the tensor."""
|
||||
return {
|
||||
'quantized': self.quantized,
|
||||
'scale': self.scale,
|
||||
'orig_dtype': self.orig_dtype
|
||||
}
|
||||
|
||||
def _load_from_state_dict(self, state_dict):
|
||||
"""Loads the tensor state from a state dictionary."""
|
||||
self.quantized = state_dict['quantized']
|
||||
self.scale = state_dict['scale']
|
||||
self.orig_dtype = state_dict['orig_dtype']
|
||||
|
||||
def __str__(self):
|
||||
return f"Auto8bitTensor({self.dequantize()})"
|
||||
|
||||
|
||||
def stochastic_grad_accummulation(param):
|
||||
if hasattr(param, "_accum_grad"):
|
||||
grad_fp32 = param._accum_grad.clone().to(torch.float32)
|
||||
grad_fp32.add_(param.grad.to(torch.float32))
|
||||
copy_stochastic(param._accum_grad, grad_fp32)
|
||||
del grad_fp32
|
||||
del param.grad
|
||||
else:
|
||||
param._accum_grad = param.grad.clone()
|
||||
del param.grad
|
||||
286
toolkit/optimizers/prodigy_8bit.py
Normal file
286
toolkit/optimizers/prodigy_8bit.py
Normal file
@@ -0,0 +1,286 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.optim import Optimizer
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, Auto8bitTensor, stochastic_grad_accummulation
|
||||
|
||||
|
||||
class Prodigy8bit(Optimizer):
|
||||
r"""
|
||||
Implements Adam with Prodigy step-sizes.
|
||||
Handles stochastic rounding for various precisions as well as stochastic gradient accumulation.
|
||||
Stores state in 8bit for memory savings.
|
||||
Leave LR set to 1 unless you encounter instability.
|
||||
|
||||
Arguments:
|
||||
params (iterable):
|
||||
Iterable of parameters to optimize or dicts defining parameter groups.
|
||||
lr (float):
|
||||
Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate.
|
||||
betas (Tuple[float, float], optional): coefficients used for computing
|
||||
running averages of gradient and its square (default: (0.9, 0.999))
|
||||
beta3 (float):
|
||||
coefficients for computing the Prodidy stepsize using running averages.
|
||||
If set to None, uses the value of square root of beta2 (default: None).
|
||||
eps (float):
|
||||
Term added to the denominator outside of the root operation to improve numerical stability. (default: 1e-8).
|
||||
weight_decay (float):
|
||||
Weight decay, i.e. a L2 penalty (default: 0).
|
||||
decouple (boolean):
|
||||
Use AdamW style decoupled weight decay
|
||||
use_bias_correction (boolean):
|
||||
Turn on Adam's bias correction. Off by default.
|
||||
safeguard_warmup (boolean):
|
||||
Remove lr from the denominator of D estimate to avoid issues during warm-up stage. Off by default.
|
||||
d0 (float):
|
||||
Initial D estimate for D-adaptation (default 1e-6). Rarely needs changing.
|
||||
d_coef (float):
|
||||
Coefficient in the expression for the estimate of d (default 1.0).
|
||||
Values such as 0.5 and 2.0 typically work as well.
|
||||
Changing this parameter is the preferred way to tune the method.
|
||||
growth_rate (float):
|
||||
prevent the D estimate from growing faster than this multiplicative rate.
|
||||
Default is inf, for unrestricted. Values like 1.02 give a kind of learning
|
||||
rate warmup effect.
|
||||
fsdp_in_use (bool):
|
||||
If you're using sharded parameters, this should be set to True. The optimizer
|
||||
will attempt to auto-detect this, but if you're using an implementation other
|
||||
than PyTorch's builtin version, the auto-detection won't work.
|
||||
"""
|
||||
|
||||
def __init__(self, params, lr=1.0,
|
||||
betas=(0.9, 0.999), beta3=None,
|
||||
eps=1e-8, weight_decay=0, decouple=True,
|
||||
use_bias_correction=False, safeguard_warmup=False,
|
||||
d0=1e-6, d_coef=1.0, growth_rate=float('inf'),
|
||||
fsdp_in_use=False):
|
||||
if not 0.0 < d0:
|
||||
raise ValueError("Invalid d0 value: {}".format(d0))
|
||||
if not 0.0 < lr:
|
||||
raise ValueError("Invalid learning rate: {}".format(lr))
|
||||
if not 0.0 < eps:
|
||||
raise ValueError("Invalid epsilon value: {}".format(eps))
|
||||
if not 0.0 <= betas[0] < 1.0:
|
||||
raise ValueError(
|
||||
"Invalid beta parameter at index 0: {}".format(betas[0]))
|
||||
if not 0.0 <= betas[1] < 1.0:
|
||||
raise ValueError(
|
||||
"Invalid beta parameter at index 1: {}".format(betas[1]))
|
||||
|
||||
if decouple and weight_decay > 0:
|
||||
print(f"Using decoupled weight decay")
|
||||
|
||||
defaults = dict(lr=lr, betas=betas, beta3=beta3,
|
||||
eps=eps, weight_decay=weight_decay,
|
||||
d=d0, d0=d0, d_max=d0,
|
||||
d_numerator=0.0, d_coef=d_coef,
|
||||
k=0, growth_rate=growth_rate,
|
||||
use_bias_correction=use_bias_correction,
|
||||
decouple=decouple, safeguard_warmup=safeguard_warmup,
|
||||
fsdp_in_use=fsdp_in_use)
|
||||
self.d0 = d0
|
||||
super(Prodigy8bit, self).__init__(params, defaults)
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# setup stochastic grad accum hooks
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_memory_efficient_fp16(self):
|
||||
return False
|
||||
|
||||
@property
|
||||
def supports_flat_params(self):
|
||||
return True
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""Performs a single optimization step.
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
# call pre step
|
||||
self.step_hook()
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
d_denom = 0.0
|
||||
|
||||
group = self.param_groups[0]
|
||||
use_bias_correction = group['use_bias_correction']
|
||||
beta1, beta2 = group['betas']
|
||||
beta3 = group['beta3']
|
||||
if beta3 is None:
|
||||
beta3 = math.sqrt(beta2)
|
||||
k = group['k']
|
||||
|
||||
d = group['d']
|
||||
d_max = group['d_max']
|
||||
d_coef = group['d_coef']
|
||||
lr = max(group['lr'] for group in self.param_groups)
|
||||
|
||||
if use_bias_correction:
|
||||
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
|
||||
else:
|
||||
bias_correction = 1
|
||||
|
||||
dlr = d*lr*bias_correction
|
||||
|
||||
growth_rate = group['growth_rate']
|
||||
decouple = group['decouple']
|
||||
fsdp_in_use = group['fsdp_in_use']
|
||||
|
||||
d_numerator = group['d_numerator']
|
||||
d_numerator *= beta3
|
||||
|
||||
for group in self.param_groups:
|
||||
decay = group['weight_decay']
|
||||
k = group['k']
|
||||
eps = group['eps']
|
||||
group_lr = group['lr']
|
||||
d0 = group['d0']
|
||||
safeguard_warmup = group['safeguard_warmup']
|
||||
|
||||
if group_lr not in [lr, 0.0]:
|
||||
raise RuntimeError(
|
||||
f"Setting different lr values in different parameter groups is only supported for values of 0")
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
if hasattr(p, "_fsdp_flattened"):
|
||||
fsdp_in_use = True
|
||||
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
# Apply weight decay (coupled variant)
|
||||
if decay != 0 and not decouple:
|
||||
grad.add_(p_fp32.data, alpha=decay)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
# State initialization
|
||||
if 'step' not in state:
|
||||
state['step'] = 0
|
||||
state['s'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
state['p0'] = Auto8bitTensor(p_fp32.detach().clone())
|
||||
# Exponential moving average of gradient values
|
||||
state['exp_avg'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
# Exponential moving average of squared gradient values
|
||||
state['exp_avg_sq'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
s = state['s'].to(torch.float32)
|
||||
p0 = state['p0'].to(torch.float32)
|
||||
|
||||
if group_lr > 0.0:
|
||||
# we use d / d0 instead of just d to avoid getting values that are too small
|
||||
d_numerator += (d / d0) * dlr * torch.dot(grad.flatten(),
|
||||
(p0.data - p_fp32.data).flatten()).item()
|
||||
|
||||
# Adam EMA updates
|
||||
exp_avg.mul_(beta1).add_(grad, alpha=d * (1-beta1))
|
||||
exp_avg_sq.mul_(beta2).addcmul_(
|
||||
grad, grad, value=d * d * (1-beta2))
|
||||
|
||||
if safeguard_warmup:
|
||||
s.mul_(beta3).add_(grad, alpha=((d / d0) * d))
|
||||
else:
|
||||
s.mul_(beta3).add_(grad, alpha=((d / d0) * dlr))
|
||||
d_denom += s.abs().sum().item()
|
||||
|
||||
# update state with stochastic rounding
|
||||
state['exp_avg'] = Auto8bitTensor(exp_avg)
|
||||
state['exp_avg_sq'] = Auto8bitTensor(exp_avg_sq)
|
||||
state['s'] = Auto8bitTensor(s)
|
||||
state['p0'] = Auto8bitTensor(p0)
|
||||
|
||||
d_hat = d
|
||||
|
||||
# if we have not done any progres, return
|
||||
# if we have any gradients available, will have d_denom > 0 (unless \|g\|=0)
|
||||
if d_denom == 0:
|
||||
return loss
|
||||
|
||||
if lr > 0.0:
|
||||
if fsdp_in_use:
|
||||
dist_tensor = torch.zeros(2).cuda()
|
||||
dist_tensor[0] = d_numerator
|
||||
dist_tensor[1] = d_denom
|
||||
dist.all_reduce(dist_tensor, op=dist.ReduceOp.SUM)
|
||||
global_d_numerator = dist_tensor[0]
|
||||
global_d_denom = dist_tensor[1]
|
||||
else:
|
||||
global_d_numerator = d_numerator
|
||||
global_d_denom = d_denom
|
||||
|
||||
d_hat = d_coef * global_d_numerator / global_d_denom
|
||||
if d == group['d0']:
|
||||
d = max(d, d_hat)
|
||||
d_max = max(d_max, d_hat)
|
||||
d = min(d_max, d * growth_rate)
|
||||
|
||||
for group in self.param_groups:
|
||||
group['d_numerator'] = global_d_numerator
|
||||
group['d_denom'] = global_d_denom
|
||||
group['d'] = d
|
||||
group['d_max'] = d_max
|
||||
group['d_hat'] = d_hat
|
||||
|
||||
decay = group['weight_decay']
|
||||
k = group['k']
|
||||
eps = group['eps']
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
state['step'] += 1
|
||||
|
||||
denom = exp_avg_sq.sqrt().add_(d * eps)
|
||||
|
||||
# Apply weight decay (decoupled variant)
|
||||
if decay != 0 and decouple:
|
||||
p_fp32.data.add_(p_fp32.data, alpha=-decay * dlr)
|
||||
|
||||
# Take step
|
||||
p_fp32.data.addcdiv_(exp_avg, denom, value=-dlr)
|
||||
# apply stochastic rounding
|
||||
copy_stochastic(p.data, p_fp32.data)
|
||||
|
||||
group['k'] = k + 1
|
||||
|
||||
return loss
|
||||
@@ -4,7 +4,7 @@ from typing import Union, List, Optional, Dict, Any, Tuple, Callable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionPipeline, LMSDiscreteScheduler, FluxPipeline
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionPipeline, LMSDiscreteScheduler, FluxPipeline, FluxControlPipeline
|
||||
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
|
||||
@@ -14,6 +14,12 @@ from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import
|
||||
from diffusers.utils import is_torch_xla_available
|
||||
from k_diffusion.external import CompVisVDenoiser, CompVisDenoiser
|
||||
from k_diffusion.sampling import get_sigmas_karras, BrownianTreeNoiseSampler
|
||||
from toolkit.models.flux import bypass_flux_guidance, restore_flux_guidance
|
||||
from diffusers.image_processor import PipelineImageInput
|
||||
from PIL import Image
|
||||
import torch.nn.functional as F
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
|
||||
if is_torch_xla_available():
|
||||
@@ -1235,6 +1241,8 @@ class FluxWithCFGPipeline(FluxPipeline):
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
# bypass the guidance embedding if there is one
|
||||
bypass_flux_guidance(self.transformer)
|
||||
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
@@ -1282,20 +1290,21 @@ class FluxWithCFGPipeline(FluxPipeline):
|
||||
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,
|
||||
)
|
||||
if guidance_scale > 1.00001:
|
||||
(
|
||||
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
|
||||
@@ -1358,21 +1367,25 @@ class FluxWithCFGPipeline(FluxPipeline):
|
||||
joint_attention_kwargs=self.joint_attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if guidance_scale > 1.00001:
|
||||
# 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]
|
||||
|
||||
# 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)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
else:
|
||||
noise_pred = noise_pred_text
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
@@ -1410,8 +1423,348 @@ class FluxWithCFGPipeline(FluxPipeline):
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
restore_flux_guidance(self.transformer)
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return FluxPipelineOutput(images=image)
|
||||
return FluxPipelineOutput(images=image)
|
||||
|
||||
|
||||
class FluxAdvancedControlPipeline(FluxControlPipeline):
|
||||
def __init__(
|
||||
self,
|
||||
scheduler,
|
||||
vae,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
text_encoder_2,
|
||||
tokenizer_2,
|
||||
transformer,
|
||||
do_inpainting=False,
|
||||
num_controls=1,
|
||||
):
|
||||
self.do_inpainting = do_inpainting
|
||||
self.num_controls = num_controls
|
||||
super().__init__(scheduler, vae, text_encoder, tokenizer, text_encoder_2, tokenizer_2, transformer)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
control_image: PipelineImageInput = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 28,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
guidance_scale: float = 3.5,
|
||||
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,
|
||||
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,
|
||||
control_image_idx: int = 0,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead
|
||||
control_image (`torch.Tensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.Tensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,:
|
||||
`List[List[torch.Tensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`):
|
||||
The ControlNet input condition to provide guidance to the `unet` for generation. If the type is
|
||||
specified as `torch.Tensor`, it is passed to ControlNet as is. `PIL.Image.Image` can also be accepted
|
||||
as an image. The dimensions of the output image defaults to `image`'s dimensions. If height and/or
|
||||
width are passed, `image` is resized accordingly. If multiple ControlNets are specified in `init`,
|
||||
images must be passed as a list such that each element of the list can be correctly batched for input
|
||||
to a single ControlNet.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, *optional*, defaults to 3.5):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
|
||||
joint_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
|
||||
is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
|
||||
images.
|
||||
"""
|
||||
|
||||
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
|
||||
|
||||
# 3. Prepare text embeddings
|
||||
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,
|
||||
)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
# num_channels_latents = self.transformer.config.in_channels // 8
|
||||
num_channels_latents = 128 // 8
|
||||
|
||||
# pull mask off control image if there is one it is a pil image
|
||||
mask = None
|
||||
if control_image is not None and self.do_inpainting and control_image.mode == "RGBA":
|
||||
control_img_array = np.array(control_image)
|
||||
mask = control_img_array[:, :, 3:4]
|
||||
# scale it to 0 - 1
|
||||
mask = mask / 255.0
|
||||
# control image ideally would be a full image here
|
||||
control_img_array = control_img_array[:, :, :3]
|
||||
control_image = Image.fromarray(control_img_array.astype(np.uint8))
|
||||
|
||||
control_image = self.prepare_image(
|
||||
image=control_image,
|
||||
width=width,
|
||||
height=height,
|
||||
batch_size=batch_size * num_images_per_prompt,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
device=device,
|
||||
dtype=self.vae.dtype,
|
||||
)
|
||||
|
||||
if control_image.ndim == 4:
|
||||
num_control_channels = num_channels_latents
|
||||
control_image = self.vae.encode(control_image).latent_dist.sample(generator=generator)
|
||||
control_image = (control_image - self.vae.config.shift_factor) * self.vae.config.scaling_factor
|
||||
|
||||
if mask is not None:
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
mask = transform(mask).to(device, dtype=control_image.dtype).unsqueeze(0)
|
||||
# resize mask to match control image
|
||||
mask = F.interpolate(mask, size=(control_image.shape[2], control_image.shape[3]), mode="bilinear", align_corners=False)
|
||||
mask = mask.to(device)
|
||||
# apply the mask to the control image so the inpaint latent area is 0
|
||||
# mask is currently 0 for inpaint area and 1 for image area
|
||||
control_image = control_image * mask
|
||||
# invert mask so it is 1 for inpaint area and 0 for image area
|
||||
mask = 1 - mask
|
||||
control_image = torch.cat([control_image, mask], dim=1)
|
||||
num_control_channels += 1
|
||||
|
||||
height_control_image, width_control_image = control_image.shape[2:]
|
||||
control_image = self._pack_latents(
|
||||
control_image,
|
||||
batch_size * num_images_per_prompt,
|
||||
num_control_channels,
|
||||
height_control_image,
|
||||
width_control_image,
|
||||
)
|
||||
|
||||
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) if sigmas is None else sigmas
|
||||
image_seq_len = latents.shape[1]
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.scheduler.config.get("base_image_seq_len", 256),
|
||||
self.scheduler.config.get("max_image_seq_len", 4096),
|
||||
self.scheduler.config.get("base_shift", 0.5),
|
||||
self.scheduler.config.get("max_shift", 1.15),
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
mu=mu,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# handle guidance
|
||||
if self.transformer.config.guidance_embeds:
|
||||
guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
|
||||
guidance = guidance.expand(latents.shape[0])
|
||||
else:
|
||||
guidance = None
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
control_image_list = []
|
||||
for idx in range(self.num_controls):
|
||||
if idx == 0 and self.do_inpainting:
|
||||
ctrl = torch.zeros_like(latents)
|
||||
# do ones for mask and zeros for image
|
||||
ctrl = torch.cat([ctrl, torch.ones_like(ctrl[:, :, :4])], dim=2)
|
||||
control_image_list.append(ctrl)
|
||||
else:
|
||||
control_image_list.append(torch.zeros_like(latents))
|
||||
|
||||
control_image_list[control_image_idx] = control_image
|
||||
|
||||
latent_model_input = torch.cat([latents] + control_image_list, dim=2)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
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]
|
||||
|
||||
# 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)
|
||||
|
||||
|
||||
31
toolkit/print.py
Normal file
31
toolkit/print.py
Normal file
@@ -0,0 +1,31 @@
|
||||
import sys
|
||||
import os
|
||||
from toolkit.accelerator import get_accelerator
|
||||
|
||||
|
||||
def print_acc(*args, **kwargs):
|
||||
if get_accelerator().is_local_main_process:
|
||||
print(*args, **kwargs)
|
||||
|
||||
|
||||
class Logger:
|
||||
def __init__(self, filename):
|
||||
self.terminal = sys.stdout
|
||||
self.log = open(filename, 'a')
|
||||
|
||||
def write(self, message):
|
||||
self.terminal.write(message)
|
||||
self.log.write(message)
|
||||
self.log.flush() # Make sure it's written immediately
|
||||
|
||||
def flush(self):
|
||||
self.terminal.flush()
|
||||
self.log.flush()
|
||||
|
||||
|
||||
def setup_log_to_file(filename):
|
||||
if get_accelerator().is_local_main_process:
|
||||
if not os.path.exists(os.path.dirname(filename)):
|
||||
os.makedirs(os.path.dirname(filename))
|
||||
sys.stdout = Logger(filename)
|
||||
sys.stderr = Logger(filename)
|
||||
@@ -62,6 +62,20 @@ class PromptEmbeds:
|
||||
prompt_embeds.attention_mask = self.attention_mask.clone()
|
||||
return prompt_embeds
|
||||
|
||||
def expand_to_batch(self, batch_size):
|
||||
pe = self.clone()
|
||||
current_batch_size = pe.text_embeds.shape[0]
|
||||
if current_batch_size == batch_size:
|
||||
return pe
|
||||
if current_batch_size != 1:
|
||||
raise Exception("Can only expand batch size for batch size 1")
|
||||
pe.text_embeds = pe.text_embeds.expand(batch_size, -1)
|
||||
if pe.pooled_embeds is not None:
|
||||
pe.pooled_embeds = pe.pooled_embeds.expand(batch_size, -1)
|
||||
if pe.attention_mask is not None:
|
||||
pe.attention_mask = pe.attention_mask.expand(batch_size, -1)
|
||||
return pe
|
||||
|
||||
|
||||
class EncodedPromptPair:
|
||||
def __init__(
|
||||
|
||||
@@ -76,6 +76,47 @@ pixart_config = {
|
||||
"variance_type": None
|
||||
}
|
||||
|
||||
flux_config = {
|
||||
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.30.0.dev0",
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": True
|
||||
}
|
||||
|
||||
sd_flow_config = {
|
||||
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.30.0.dev0",
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 3.0,
|
||||
"use_dynamic_shifting": False
|
||||
}
|
||||
|
||||
lumina2_config = {
|
||||
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
||||
"_diffusers_version": "0.33.0.dev0",
|
||||
"base_image_seq_len": 256,
|
||||
"base_shift": 0.5,
|
||||
"invert_sigmas": False,
|
||||
"max_image_seq_len": 4096,
|
||||
"max_shift": 1.15,
|
||||
"num_train_timesteps": 1000,
|
||||
"shift": 6.0,
|
||||
"shift_terminal": None,
|
||||
"use_beta_sigmas": False,
|
||||
"use_dynamic_shifting": False,
|
||||
"use_exponential_sigmas": False,
|
||||
"use_karras_sigmas": False
|
||||
}
|
||||
|
||||
|
||||
def get_sampler(
|
||||
sampler: str,
|
||||
@@ -120,12 +161,16 @@ def get_sampler(
|
||||
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
|
||||
}
|
||||
config_to_use = copy.deepcopy(flux_config)
|
||||
if arch == "sd":
|
||||
config_to_use = copy.deepcopy(sd_flow_config)
|
||||
if arch == "flux":
|
||||
config_to_use = copy.deepcopy(flux_config)
|
||||
elif arch == "lumina2":
|
||||
config_to_use = copy.deepcopy(lumina2_config)
|
||||
else:
|
||||
# use flux by default
|
||||
config_to_use = copy.deepcopy(flux_config)
|
||||
else:
|
||||
raise ValueError(f"Sampler {sampler} not supported")
|
||||
|
||||
|
||||
@@ -1,14 +1,29 @@
|
||||
import math
|
||||
from typing import Union
|
||||
|
||||
from torch.distributions import LogNormal
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len,
|
||||
base_seq_len: int = 256,
|
||||
max_seq_len: int = 4096,
|
||||
base_shift: float = 0.5,
|
||||
max_shift: float = 1.16,
|
||||
):
|
||||
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
b = base_shift - m * base_seq_len
|
||||
mu = image_seq_len * m + b
|
||||
return mu
|
||||
|
||||
|
||||
class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.init_noise_sigma = 1.0
|
||||
self.timestep_type = "linear"
|
||||
|
||||
with torch.no_grad():
|
||||
# create weights for timesteps
|
||||
@@ -25,19 +40,31 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
# Scale to make mean 1
|
||||
bsmntw_weighing = y_shifted * (num_timesteps / y_shifted.sum())
|
||||
|
||||
# only do half bell
|
||||
hbsmntw_weighing = y_shifted * (num_timesteps / y_shifted.sum())
|
||||
|
||||
# flatten second half to max
|
||||
hbsmntw_weighing[num_timesteps //
|
||||
2:] = hbsmntw_weighing[num_timesteps // 2:].max()
|
||||
|
||||
# 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
|
||||
self.linear_timesteps_weights2 = hbsmntw_weighing
|
||||
pass
|
||||
|
||||
def get_weights_for_timesteps(self, timesteps: torch.Tensor) -> torch.Tensor:
|
||||
def get_weights_for_timesteps(self, timesteps: torch.Tensor, v2=False) -> torch.Tensor:
|
||||
# Get the indices of the timesteps
|
||||
step_indices = [(self.timesteps == t).nonzero().item() for t in 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()
|
||||
if v2:
|
||||
weights = self.linear_timesteps_weights2[step_indices].flatten()
|
||||
else:
|
||||
weights = self.linear_timesteps_weights[step_indices].flatten()
|
||||
|
||||
return weights
|
||||
|
||||
@@ -45,7 +72,8 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
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]
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
@@ -59,32 +87,30 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
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
|
||||
# forward ODE
|
||||
noisy_model_input = (1.0 - t_01) * original_samples + t_01 * noise
|
||||
# reverse ODE
|
||||
# noisy_model_input = (1 - t_01) * noise + t_01 * original_samples
|
||||
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:
|
||||
def set_train_timesteps(
|
||||
self,
|
||||
num_timesteps,
|
||||
device,
|
||||
timestep_type='linear',
|
||||
latents=None,
|
||||
patch_size=1
|
||||
):
|
||||
self.timestep_type = timestep_type
|
||||
if timestep_type == 'linear':
|
||||
timesteps = torch.linspace(1000, 0, num_timesteps, device=device)
|
||||
self.timesteps = timesteps
|
||||
return timesteps
|
||||
else:
|
||||
elif timestep_type == 'sigmoid':
|
||||
# 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))
|
||||
@@ -98,3 +124,89 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
|
||||
return timesteps
|
||||
elif timestep_type in ['flux_shift', 'lumina2_shift', 'shift']:
|
||||
# matches inference dynamic shifting
|
||||
timesteps = np.linspace(
|
||||
self._sigma_to_t(self.sigma_max), self._sigma_to_t(
|
||||
self.sigma_min), num_timesteps
|
||||
)
|
||||
|
||||
sigmas = timesteps / self.config.num_train_timesteps
|
||||
|
||||
if self.config.use_dynamic_shifting:
|
||||
if latents is None:
|
||||
raise ValueError('latents is None')
|
||||
|
||||
# for flux we double up the patch size before sending her to simulate the latent reduction
|
||||
h = latents.shape[2]
|
||||
w = latents.shape[3]
|
||||
image_seq_len = h * w // (patch_size**2)
|
||||
|
||||
mu = calculate_shift(
|
||||
image_seq_len,
|
||||
self.config.get("base_image_seq_len", 256),
|
||||
self.config.get("max_image_seq_len", 4096),
|
||||
self.config.get("base_shift", 0.5),
|
||||
self.config.get("max_shift", 1.16),
|
||||
)
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas)
|
||||
else:
|
||||
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
|
||||
|
||||
if self.config.shift_terminal:
|
||||
sigmas = self.stretch_shift_to_terminal(sigmas)
|
||||
|
||||
if self.config.use_karras_sigmas:
|
||||
sigmas = self._convert_to_karras(
|
||||
in_sigmas=sigmas, num_inference_steps=self.config.num_train_timesteps)
|
||||
elif self.config.use_exponential_sigmas:
|
||||
sigmas = self._convert_to_exponential(
|
||||
in_sigmas=sigmas, num_inference_steps=self.config.num_train_timesteps)
|
||||
elif self.config.use_beta_sigmas:
|
||||
sigmas = self._convert_to_beta(
|
||||
in_sigmas=sigmas, num_inference_steps=self.config.num_train_timesteps)
|
||||
|
||||
sigmas = torch.from_numpy(sigmas).to(
|
||||
dtype=torch.float32, device=device)
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
|
||||
if self.config.invert_sigmas:
|
||||
sigmas = 1.0 - sigmas
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
sigmas = torch.cat(
|
||||
[sigmas, torch.ones(1, device=sigmas.device)])
|
||||
else:
|
||||
sigmas = torch.cat(
|
||||
[sigmas, torch.zeros(1, device=sigmas.device)])
|
||||
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
self.sigmas = sigmas
|
||||
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
return timesteps
|
||||
|
||||
elif timestep_type == 'lognorm_blend':
|
||||
# disgtribute timestepd to the center/early and blend in linear
|
||||
alpha = 0.75
|
||||
|
||||
lognormal = LogNormal(loc=0, scale=0.333)
|
||||
|
||||
# Sample from the distribution
|
||||
t1 = lognormal.sample((int(num_timesteps * alpha),)).to(device)
|
||||
|
||||
# Scale and reverse the values to go from 1000 to 0
|
||||
t1 = ((1 - t1/t1.max()) * 1000)
|
||||
|
||||
# add half of linear
|
||||
t2 = torch.linspace(1000, 0, int(
|
||||
num_timesteps * (1 - alpha)), device=device)
|
||||
timesteps = torch.cat((t1, t2))
|
||||
|
||||
# Sort the timesteps in descending order
|
||||
timesteps, _ = torch.sort(timesteps, descending=True)
|
||||
|
||||
timesteps = timesteps.to(torch.int)
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
return timesteps
|
||||
else:
|
||||
raise ValueError(f"Invalid timestep type: {timestep_type}")
|
||||
|
||||
@@ -39,7 +39,10 @@ def get_train_sd_device_state_preset(
|
||||
train_lora: bool = False,
|
||||
train_adapter: bool = False,
|
||||
train_embedding: bool = False,
|
||||
train_decorator: bool = False,
|
||||
train_refiner: bool = False,
|
||||
unload_text_encoder: bool = False,
|
||||
require_grads: bool = True,
|
||||
):
|
||||
preset = copy.deepcopy(empty_preset)
|
||||
if not cached_latents:
|
||||
@@ -47,27 +50,27 @@ def get_train_sd_device_state_preset(
|
||||
|
||||
if train_unet:
|
||||
preset['unet']['training'] = True
|
||||
preset['unet']['requires_grad'] = True
|
||||
preset['unet']['requires_grad'] = require_grads
|
||||
preset['unet']['device'] = device
|
||||
else:
|
||||
preset['unet']['device'] = device
|
||||
|
||||
if train_text_encoder:
|
||||
preset['text_encoder']['training'] = True
|
||||
preset['text_encoder']['requires_grad'] = True
|
||||
preset['text_encoder']['requires_grad'] = require_grads
|
||||
preset['text_encoder']['device'] = device
|
||||
else:
|
||||
preset['text_encoder']['device'] = device
|
||||
|
||||
if train_embedding:
|
||||
preset['text_encoder']['training'] = True
|
||||
preset['text_encoder']['requires_grad'] = True
|
||||
preset['text_encoder']['requires_grad'] = require_grads
|
||||
preset['text_encoder']['training'] = True
|
||||
preset['unet']['training'] = True
|
||||
|
||||
if train_refiner:
|
||||
preset['refiner_unet']['training'] = True
|
||||
preset['refiner_unet']['requires_grad'] = True
|
||||
preset['refiner_unet']['requires_grad'] = require_grads
|
||||
preset['refiner_unet']['device'] = device
|
||||
# if not training unet, move that to cpu
|
||||
if not train_unet:
|
||||
@@ -80,12 +83,25 @@ def get_train_sd_device_state_preset(
|
||||
preset['refiner_unet']['requires_grad'] = False
|
||||
|
||||
if train_adapter:
|
||||
preset['adapter']['requires_grad'] = True
|
||||
preset['adapter']['requires_grad'] = require_grads
|
||||
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
|
||||
|
||||
if train_decorator:
|
||||
preset['text_encoder']['training'] = False
|
||||
preset['text_encoder']['requires_grad'] = False
|
||||
preset['text_encoder']['device'] = device
|
||||
preset['unet']['training'] = True
|
||||
preset['unet']['requires_grad'] = False
|
||||
preset['unet']['device'] = device
|
||||
|
||||
if unload_text_encoder:
|
||||
preset['text_encoder']['training'] = False
|
||||
preset['text_encoder']['requires_grad'] = False
|
||||
preset['text_encoder']['device'] = 'cpu'
|
||||
|
||||
return preset
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,6 +1,10 @@
|
||||
import time
|
||||
from collections import OrderedDict, deque
|
||||
import sys
|
||||
import os
|
||||
|
||||
# check if is ui process will have IS_AI_TOOLKIT_UI in env
|
||||
is_ui = os.environ.get("IS_AI_TOOLKIT_UI", "0") == "1"
|
||||
|
||||
class Timer:
|
||||
def __init__(self, name='Timer', max_buffer=10):
|
||||
@@ -9,6 +13,7 @@ class Timer:
|
||||
self.timers = OrderedDict()
|
||||
self.active_timers = {}
|
||||
self.current_timer = None # Used for the context manager functionality
|
||||
self._after_print_hooks = []
|
||||
|
||||
def start(self, timer_name):
|
||||
if timer_name not in self.timers:
|
||||
@@ -34,14 +39,25 @@ class Timer:
|
||||
if len(self.timers[timer_name]) > self.max_buffer:
|
||||
self.timers[timer_name].popleft()
|
||||
|
||||
def add_after_print_hook(self, hook):
|
||||
self._after_print_hooks.append(hook)
|
||||
|
||||
def print(self):
|
||||
print(f"\nTimer '{self.name}':")
|
||||
if not is_ui:
|
||||
print(f"\nTimer '{self.name}':")
|
||||
timing_dict = {}
|
||||
# sort by longest at top
|
||||
for timer_name, timings in sorted(self.timers.items(), key=lambda x: sum(x[1]), reverse=True):
|
||||
avg_time = sum(timings) / len(timings)
|
||||
print(f" - {avg_time:.4f}s avg - {timer_name}, num = {len(timings)}")
|
||||
|
||||
if not is_ui:
|
||||
print(f" - {avg_time:.4f}s avg - {timer_name}, num = {len(timings)}")
|
||||
timing_dict[timer_name] = avg_time
|
||||
|
||||
print('')
|
||||
for hook in self._after_print_hooks:
|
||||
hook(timing_dict)
|
||||
if not is_ui:
|
||||
print('')
|
||||
|
||||
def reset(self):
|
||||
self.timers.clear()
|
||||
|
||||
@@ -137,6 +137,8 @@ def match_noise_to_target_mean_offset(noise, target, mix=0.5, dim=None):
|
||||
def apply_noise_offset(noise, noise_offset):
|
||||
if noise_offset is None or (noise_offset < 0.000001 and noise_offset > -0.000001):
|
||||
return noise
|
||||
if len(noise.shape) > 4:
|
||||
raise ValueError("Applying noise offset not supported for video models at this time.")
|
||||
noise = noise + noise_offset * torch.randn((noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
|
||||
return noise
|
||||
|
||||
@@ -517,6 +519,7 @@ def encode_prompts_flux(
|
||||
truncate: bool = True,
|
||||
max_length=None,
|
||||
dropout_prob=0.0,
|
||||
attn_mask: bool = False,
|
||||
):
|
||||
if max_length is None:
|
||||
max_length = 512
|
||||
@@ -568,12 +571,9 @@ def encode_prompts_flux(
|
||||
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)
|
||||
if attn_mask:
|
||||
prompt_attention_mask = text_inputs["attention_mask"].unsqueeze(-1).expand(prompt_embeds.shape)
|
||||
prompt_embeds = prompt_embeds * prompt_attention_mask.to(dtype=prompt_embeds.dtype, device=prompt_embeds.device)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds
|
||||
|
||||
|
||||
@@ -1,120 +0,0 @@
|
||||
# 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
|
||||
12
toolkit/util/get_model.py
Normal file
12
toolkit/util/get_model.py
Normal file
@@ -0,0 +1,12 @@
|
||||
from toolkit.stable_diffusion_model import StableDiffusion
|
||||
from toolkit.config_modules import ModelConfig
|
||||
|
||||
def get_model_class(config: ModelConfig):
|
||||
if config.arch == "wan21":
|
||||
from toolkit.models.wan21 import Wan21
|
||||
return Wan21
|
||||
elif config.arch == "cogview4":
|
||||
from toolkit.models.cogview4 import CogView4
|
||||
return CogView4
|
||||
else:
|
||||
return StableDiffusion
|
||||
226
toolkit/util/mask.py
Normal file
226
toolkit/util/mask.py
Normal file
@@ -0,0 +1,226 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import os
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
import time
|
||||
|
||||
|
||||
def generate_random_mask(
|
||||
batch_size,
|
||||
height=256,
|
||||
width=256,
|
||||
device='cuda',
|
||||
min_coverage=0.2,
|
||||
max_coverage=0.8,
|
||||
num_blobs_range=(1, 3)
|
||||
):
|
||||
"""
|
||||
Generate random blob masks for a batch of images.
|
||||
Fast GPU version with smooth, non-circular blob shapes.
|
||||
|
||||
Args:
|
||||
batch_size (int): Number of masks to generate
|
||||
height (int): Height of the mask
|
||||
width (int): Width of the mask
|
||||
device (str): Device to run the computation on ('cuda' or 'cpu')
|
||||
min_coverage (float): Minimum percentage of the image to be covered (0-1)
|
||||
max_coverage (float): Maximum percentage of the image to be covered (0-1)
|
||||
num_blobs_range (tuple): Range of number of blobs (min, max)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Binary masks with shape (batch_size, 1, height, width)
|
||||
"""
|
||||
# Initialize masks on GPU
|
||||
masks = torch.zeros((batch_size, 1, height, width), device=device)
|
||||
|
||||
# Pre-compute coordinate grid on GPU
|
||||
y_indices = torch.arange(height, device=device).view(
|
||||
height, 1).expand(height, width)
|
||||
x_indices = torch.arange(width, device=device).view(
|
||||
1, width).expand(height, width)
|
||||
|
||||
# Prepare gaussian kernels for smoothing
|
||||
small_kernel = get_gaussian_kernel(7, 1.0).to(device)
|
||||
small_kernel = small_kernel.view(1, 1, 7, 7)
|
||||
|
||||
large_kernel = get_gaussian_kernel(15, 2.5).to(device)
|
||||
large_kernel = large_kernel.view(1, 1, 15, 15)
|
||||
|
||||
# Constants
|
||||
max_radius = min(height, width) // 3
|
||||
min_radius = min(height, width) // 8
|
||||
|
||||
# For each mask in the batch
|
||||
for b in range(batch_size):
|
||||
# Determine number of blobs for this mask
|
||||
num_blobs = np.random.randint(
|
||||
num_blobs_range[0], num_blobs_range[1] + 1)
|
||||
|
||||
# Target coverage for this mask
|
||||
target_coverage = np.random.uniform(min_coverage, max_coverage)
|
||||
|
||||
# Initialize this mask
|
||||
mask = torch.zeros(1, 1, height, width, device=device)
|
||||
|
||||
# Generate blobs with smoother edges
|
||||
for _ in range(num_blobs):
|
||||
# Create a low-frequency noise field first (for smooth organic shapes)
|
||||
noise_field = torch.zeros(height, width, device=device)
|
||||
|
||||
# Use low-frequency sine waves to create base shape distortion
|
||||
# This creates smoother warping compared to pure random noise
|
||||
num_waves = np.random.randint(2, 5)
|
||||
for i in range(num_waves):
|
||||
freq_x = np.random.uniform(1.0, 3.0) * np.pi / width
|
||||
freq_y = np.random.uniform(1.0, 3.0) * np.pi / height
|
||||
phase_x = np.random.uniform(0, 2 * np.pi)
|
||||
phase_y = np.random.uniform(0, 2 * np.pi)
|
||||
amp = np.random.uniform(0.5, 1.0) * max_radius / (i+1.5)
|
||||
|
||||
# Generate smooth wave patterns
|
||||
wave = torch.sin(x_indices * freq_x + phase_x) * \
|
||||
torch.sin(y_indices * freq_y + phase_y) * amp
|
||||
noise_field += wave
|
||||
|
||||
# Basic ellipse parameters
|
||||
center_y = np.random.randint(height//4, 3*height//4)
|
||||
center_x = np.random.randint(width//4, 3*width//4)
|
||||
radius = np.random.randint(min_radius, max_radius)
|
||||
|
||||
# Squeeze and stretch the ellipse with random scaling
|
||||
scale_y = np.random.uniform(0.6, 1.4)
|
||||
scale_x = np.random.uniform(0.6, 1.4)
|
||||
|
||||
# Random rotation
|
||||
theta = np.random.uniform(0, 2 * np.pi)
|
||||
cos_theta, sin_theta = np.cos(theta), np.sin(theta)
|
||||
|
||||
# Calculate elliptical distance field
|
||||
y_scaled = (y_indices - center_y) * scale_y
|
||||
x_scaled = (x_indices - center_x) * scale_x
|
||||
|
||||
# Apply rotation
|
||||
rotated_y = y_scaled * cos_theta - x_scaled * sin_theta
|
||||
rotated_x = y_scaled * sin_theta + x_scaled * cos_theta
|
||||
|
||||
# Compute distances
|
||||
distances = torch.sqrt(rotated_y**2 + rotated_x**2)
|
||||
|
||||
# Apply the smooth noise field to the distance field
|
||||
perturbed_distances = distances + noise_field
|
||||
|
||||
# Create base blob
|
||||
blob = (perturbed_distances < radius).float(
|
||||
).unsqueeze(0).unsqueeze(0)
|
||||
|
||||
# Apply strong smoothing for very smooth edges
|
||||
# Double smoothing to get really organic edges
|
||||
blob = F.pad(blob, (7, 7, 7, 7), mode='reflect')
|
||||
blob = F.conv2d(blob, large_kernel, padding=0)
|
||||
|
||||
# Apply threshold to get a nice shape
|
||||
rand_threshold = np.random.uniform(0.3, 0.6)
|
||||
blob = (blob > rand_threshold).float()
|
||||
|
||||
# Apply second smoothing pass
|
||||
blob = F.pad(blob, (3, 3, 3, 3), mode='reflect')
|
||||
blob = F.conv2d(blob, small_kernel, padding=0)
|
||||
blob = (blob > 0.5).float()
|
||||
|
||||
# Add to mask
|
||||
mask = torch.maximum(mask, blob)
|
||||
|
||||
# Ensure desired coverage
|
||||
current_coverage = mask.mean().item()
|
||||
|
||||
# Scale if needed to match target coverage
|
||||
if current_coverage > 0: # Avoid division by zero
|
||||
if current_coverage < target_coverage * 0.7: # Too small
|
||||
# Dilate mask to increase coverage
|
||||
mask = F.pad(mask, (2, 2, 2, 2), mode='reflect')
|
||||
mask = F.max_pool2d(mask, kernel_size=5, stride=1, padding=0)
|
||||
elif current_coverage > target_coverage * 1.3: # Too large
|
||||
# Erode mask to decrease coverage
|
||||
mask = F.pad(mask, (1, 1, 1, 1), mode='reflect')
|
||||
mask = F.avg_pool2d(mask, kernel_size=3, stride=1, padding=0)
|
||||
mask = (mask > 0.7).float()
|
||||
|
||||
# Final smooth and threshold
|
||||
mask = F.pad(mask, (3, 3, 3, 3), mode='reflect')
|
||||
mask = F.conv2d(mask, small_kernel, padding=0)
|
||||
mask = (mask > 0.5).float()
|
||||
|
||||
# Add to batch
|
||||
masks[b] = mask
|
||||
|
||||
return masks
|
||||
|
||||
|
||||
def get_gaussian_kernel(kernel_size=5, sigma=1.0):
|
||||
"""
|
||||
Returns a 2D Gaussian kernel.
|
||||
"""
|
||||
# Create 1D kernels
|
||||
x = torch.linspace(-sigma * 2, sigma * 2, kernel_size)
|
||||
x = x.view(1, -1).repeat(kernel_size, 1)
|
||||
y = x.transpose(0, 1)
|
||||
|
||||
# 2D Gaussian
|
||||
gaussian = torch.exp(-(x**2 + y**2) / (2 * sigma**2))
|
||||
gaussian /= gaussian.sum()
|
||||
|
||||
return gaussian
|
||||
|
||||
|
||||
def save_masks_as_images(masks, output_dir="output"):
|
||||
"""
|
||||
Save generated masks as RGB JPG images using PIL.
|
||||
"""
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
batch_size = masks.shape[0]
|
||||
for i in range(batch_size):
|
||||
# Convert mask to numpy array
|
||||
mask = masks[i, 0].cpu().numpy()
|
||||
|
||||
# Scale to 0-255 range and convert to uint8
|
||||
mask_255 = (mask * 255).astype(np.uint8)
|
||||
|
||||
# Create RGB image (white mask on black background)
|
||||
rgb_mask = np.stack([mask_255, mask_255, mask_255], axis=2)
|
||||
|
||||
# Convert to PIL Image and save
|
||||
img = Image.fromarray(rgb_mask)
|
||||
img.save(os.path.join(output_dir, f"mask_{i:03d}.jpg"), quality=95)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Parameters
|
||||
batch_size = 20
|
||||
height = 256
|
||||
width = 256
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
|
||||
print(f"Generating {batch_size} random blob masks on {device}...")
|
||||
|
||||
for i in range(5):
|
||||
# time it
|
||||
start = time.time()
|
||||
masks = generate_random_mask(
|
||||
batch_size=batch_size,
|
||||
height=height,
|
||||
width=width,
|
||||
device=device,
|
||||
min_coverage=0.2,
|
||||
max_coverage=0.8,
|
||||
num_blobs_range=(1, 3)
|
||||
)
|
||||
end = time.time()
|
||||
# print time in milliseconds
|
||||
print(f"Time taken: {(end - start)*1000:.2f} ms")
|
||||
|
||||
print(f"Saving masks to 'output' directory...")
|
||||
save_masks_as_images(masks)
|
||||
|
||||
print("Done!")
|
||||
98
toolkit/util/quantize.py
Normal file
98
toolkit/util/quantize.py
Normal file
@@ -0,0 +1,98 @@
|
||||
from fnmatch import fnmatch
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
|
||||
from optimum.quanto.quantize import _quantize_submodule
|
||||
from optimum.quanto.tensor import Optimizer, qtype, qtypes
|
||||
from torchao.quantization.quant_api import (
|
||||
quantize_ as torchao_quantize_,
|
||||
Float8WeightOnlyConfig,
|
||||
UIntXWeightOnlyConfig
|
||||
)
|
||||
|
||||
# the quantize function in quanto had a bug where it was using exclude instead of include
|
||||
|
||||
Q_MODULES = ['QLinear', 'QConv2d', 'QEmbedding', 'QBatchNorm2d', 'QLayerNorm', 'QConvTranspose2d', 'QEmbeddingBag']
|
||||
|
||||
torchao_qtypes = {
|
||||
# "int4": Int4WeightOnlyConfig(),
|
||||
"uint2": UIntXWeightOnlyConfig(torch.uint2),
|
||||
"uint3": UIntXWeightOnlyConfig(torch.uint3),
|
||||
"uint4": UIntXWeightOnlyConfig(torch.uint4),
|
||||
"uint5": UIntXWeightOnlyConfig(torch.uint5),
|
||||
"uint6": UIntXWeightOnlyConfig(torch.uint6),
|
||||
"uint7": UIntXWeightOnlyConfig(torch.uint7),
|
||||
"uint8": UIntXWeightOnlyConfig(torch.uint8),
|
||||
"float8": Float8WeightOnlyConfig(),
|
||||
}
|
||||
|
||||
class aotype:
|
||||
def __init__(self, name: str):
|
||||
self.name = name
|
||||
self.config = torchao_qtypes[name]
|
||||
|
||||
def get_qtype(qtype: Union[str, qtype]) -> qtype:
|
||||
if qtype in torchao_qtypes:
|
||||
return aotype(qtype)
|
||||
if isinstance(qtype, str):
|
||||
return qtypes[qtype]
|
||||
else:
|
||||
return qtype
|
||||
|
||||
def quantize(
|
||||
model: torch.nn.Module,
|
||||
weights: Optional[Union[str, qtype, aotype]] = None,
|
||||
activations: Optional[Union[str, qtype]] = None,
|
||||
optimizer: Optional[Optimizer] = None,
|
||||
include: Optional[Union[str, List[str]]] = None,
|
||||
exclude: Optional[Union[str, List[str]]] = None,
|
||||
):
|
||||
"""Quantize the specified model submodules
|
||||
|
||||
Recursively quantize the submodules of the specified parent model.
|
||||
|
||||
Only modules that have quantized counterparts will be quantized.
|
||||
|
||||
If include patterns are specified, the submodule name must match one of them.
|
||||
|
||||
If exclude patterns are specified, the submodule must not match one of them.
|
||||
|
||||
Include or exclude patterns are Unix shell-style wildcards which are NOT regular expressions. See
|
||||
https://docs.python.org/3/library/fnmatch.html for more details.
|
||||
|
||||
Note: quantization happens in-place and modifies the original model and its descendants.
|
||||
|
||||
Args:
|
||||
model (`torch.nn.Module`): the model whose submodules will be quantized.
|
||||
weights (`Optional[Union[str, qtype]]`): the qtype for weights quantization.
|
||||
activations (`Optional[Union[str, qtype]]`): the qtype for activations quantization.
|
||||
include (`Optional[Union[str, List[str]]]`):
|
||||
Patterns constituting the allowlist. If provided, module names must match at
|
||||
least one pattern from the allowlist.
|
||||
exclude (`Optional[Union[str, List[str]]]`):
|
||||
Patterns constituting the denylist. If provided, module names must not match
|
||||
any patterns from the denylist.
|
||||
"""
|
||||
if include is not None:
|
||||
include = [include] if isinstance(include, str) else include
|
||||
if exclude is not None:
|
||||
exclude = [exclude] if isinstance(exclude, str) else exclude
|
||||
for name, m in model.named_modules():
|
||||
if include is not None and not any(fnmatch(name, pattern) for pattern in include):
|
||||
continue
|
||||
if exclude is not None and any(fnmatch(name, pattern) for pattern in exclude):
|
||||
continue
|
||||
try:
|
||||
# check if m is QLinear or QConv2d
|
||||
if m.__class__.__name__ in Q_MODULES:
|
||||
continue
|
||||
else:
|
||||
if isinstance(weights, aotype):
|
||||
torchao_quantize_(m, weights.config)
|
||||
else:
|
||||
_quantize_submodule(model, name, m, weights=weights,
|
||||
activations=activations, optimizer=optimizer)
|
||||
except Exception as e:
|
||||
print(f"Failed to quantize {name}: {e}")
|
||||
raise e
|
||||
39
toolkit/util/wavelet_loss.py
Normal file
39
toolkit/util/wavelet_loss.py
Normal file
@@ -0,0 +1,39 @@
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
_dwt = None
|
||||
|
||||
|
||||
def _get_wavelet_loss(device, dtype):
|
||||
global _dwt
|
||||
if _dwt is not None:
|
||||
return _dwt
|
||||
|
||||
# init wavelets
|
||||
from pytorch_wavelets import DWTForward
|
||||
# wave='db1' wave='haar'
|
||||
dwt = DWTForward(J=1, mode='zero', wave='haar').to(
|
||||
device=device, dtype=dtype)
|
||||
_dwt = dwt
|
||||
return dwt
|
||||
|
||||
|
||||
def wavelet_loss(model_pred, latents, noise):
|
||||
model_pred = model_pred.float()
|
||||
latents = latents.float()
|
||||
noise = noise.float()
|
||||
dwt = _get_wavelet_loss(model_pred.device, model_pred.dtype)
|
||||
with torch.no_grad():
|
||||
model_input_xll, model_input_xh = dwt(latents)
|
||||
model_input_xlh, model_input_xhl, model_input_xhh = torch.unbind(model_input_xh[0], dim=2)
|
||||
model_input = torch.cat([model_input_xll, model_input_xlh, model_input_xhl, model_input_xhh], dim=1)
|
||||
|
||||
# reverse the noise to get the model prediction of the pure latents
|
||||
model_pred = noise - model_pred
|
||||
|
||||
model_pred_xll, model_pred_xh = dwt(model_pred)
|
||||
model_pred_xlh, model_pred_xhl, model_pred_xhh = torch.unbind(model_pred_xh[0], dim=2)
|
||||
model_pred = torch.cat([model_pred_xll, model_pred_xlh, model_pred_xhl, model_pred_xhh], dim=1)
|
||||
|
||||
return torch.nn.functional.mse_loss(model_pred, model_input, reduction="none")
|
||||
42
ui/.gitignore
vendored
Normal file
42
ui/.gitignore
vendored
Normal file
@@ -0,0 +1,42 @@
|
||||
# See https://help.github.com/articles/ignoring-files/ for more about ignoring files.
|
||||
|
||||
# dependencies
|
||||
/node_modules
|
||||
/.pnp
|
||||
.pnp.*
|
||||
.yarn/*
|
||||
!.yarn/patches
|
||||
!.yarn/plugins
|
||||
!.yarn/releases
|
||||
!.yarn/versions
|
||||
|
||||
# testing
|
||||
/coverage
|
||||
|
||||
# next.js
|
||||
/.next/
|
||||
/out/
|
||||
|
||||
# production
|
||||
/build
|
||||
|
||||
# misc
|
||||
.DS_Store
|
||||
*.pem
|
||||
|
||||
# debug
|
||||
npm-debug.log*
|
||||
yarn-debug.log*
|
||||
yarn-error.log*
|
||||
.pnpm-debug.log*
|
||||
|
||||
# env files (can opt-in for committing if needed)
|
||||
.env*
|
||||
|
||||
# vercel
|
||||
.vercel
|
||||
|
||||
# typescript
|
||||
*.tsbuildinfo
|
||||
next-env.d.ts
|
||||
aitk_db.db
|
||||
36
ui/README.md
Normal file
36
ui/README.md
Normal file
@@ -0,0 +1,36 @@
|
||||
This is a [Next.js](https://nextjs.org) project bootstrapped with [`create-next-app`](https://nextjs.org/docs/app/api-reference/cli/create-next-app).
|
||||
|
||||
## Getting Started
|
||||
|
||||
First, run the development server:
|
||||
|
||||
```bash
|
||||
npm run dev
|
||||
# or
|
||||
yarn dev
|
||||
# or
|
||||
pnpm dev
|
||||
# or
|
||||
bun dev
|
||||
```
|
||||
|
||||
Open [http://localhost:3000](http://localhost:3000) with your browser to see the result.
|
||||
|
||||
You can start editing the page by modifying `app/page.tsx`. The page auto-updates as you edit the file.
|
||||
|
||||
This project uses [`next/font`](https://nextjs.org/docs/app/building-your-application/optimizing/fonts) to automatically optimize and load [Geist](https://vercel.com/font), a new font family for Vercel.
|
||||
|
||||
## Learn More
|
||||
|
||||
To learn more about Next.js, take a look at the following resources:
|
||||
|
||||
- [Next.js Documentation](https://nextjs.org/docs) - learn about Next.js features and API.
|
||||
- [Learn Next.js](https://nextjs.org/learn) - an interactive Next.js tutorial.
|
||||
|
||||
You can check out [the Next.js GitHub repository](https://github.com/vercel/next.js) - your feedback and contributions are welcome!
|
||||
|
||||
## Deploy on Vercel
|
||||
|
||||
The easiest way to deploy your Next.js app is to use the [Vercel Platform](https://vercel.com/new?utm_medium=default-template&filter=next.js&utm_source=create-next-app&utm_campaign=create-next-app-readme) from the creators of Next.js.
|
||||
|
||||
Check out our [Next.js deployment documentation](https://nextjs.org/docs/app/building-your-application/deploying) for more details.
|
||||
15
ui/next.config.ts
Normal file
15
ui/next.config.ts
Normal file
@@ -0,0 +1,15 @@
|
||||
import type { NextConfig } from 'next';
|
||||
|
||||
const nextConfig: NextConfig = {
|
||||
typescript: {
|
||||
// Remove this. Build fails because of route types
|
||||
ignoreBuildErrors: true,
|
||||
},
|
||||
experimental: {
|
||||
serverActions: {
|
||||
bodySizeLimit: '100mb',
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
export default nextConfig;
|
||||
4549
ui/package-lock.json
generated
Normal file
4549
ui/package-lock.json
generated
Normal file
File diff suppressed because it is too large
Load Diff
43
ui/package.json
Normal file
43
ui/package.json
Normal file
@@ -0,0 +1,43 @@
|
||||
{
|
||||
"name": "ai-toolkit-ui",
|
||||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"scripts": {
|
||||
"dev": "next dev --turbopack",
|
||||
"build": "next build",
|
||||
"start": "next start --port 8675",
|
||||
"build_and_start": "npm install && npm run update_db && npm run build && npm run start",
|
||||
"lint": "next lint",
|
||||
"update_db": "npx prisma generate && npx prisma db push",
|
||||
"format": "prettier --write \"**/*.{js,jsx,ts,tsx,css,scss}\""
|
||||
},
|
||||
"dependencies": {
|
||||
"@headlessui/react": "^2.2.0",
|
||||
"@monaco-editor/react": "^4.7.0",
|
||||
"@prisma/client": "^6.3.1",
|
||||
"axios": "^1.7.9",
|
||||
"classnames": "^2.5.1",
|
||||
"lucide-react": "^0.475.0",
|
||||
"next": "15.1.7",
|
||||
"node-cache": "^5.1.2",
|
||||
"prisma": "^6.3.1",
|
||||
"react": "^19.0.0",
|
||||
"react-dom": "^19.0.0",
|
||||
"react-dropzone": "^14.3.5",
|
||||
"react-global-hooks": "^1.3.5",
|
||||
"react-icons": "^5.5.0",
|
||||
"sqlite3": "^5.1.7",
|
||||
"yaml": "^2.7.0"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/node": "^20",
|
||||
"@types/react": "^19",
|
||||
"@types/react-dom": "^19",
|
||||
"postcss": "^8",
|
||||
"prettier": "^3.5.1",
|
||||
"prettier-basic": "^1.0.0",
|
||||
"tailwindcss": "^3.4.1",
|
||||
"typescript": "^5"
|
||||
},
|
||||
"prettier": "prettier-basic"
|
||||
}
|
||||
8
ui/postcss.config.mjs
Normal file
8
ui/postcss.config.mjs
Normal file
@@ -0,0 +1,8 @@
|
||||
/** @type {import('postcss-load-config').Config} */
|
||||
const config = {
|
||||
plugins: {
|
||||
tailwindcss: {},
|
||||
},
|
||||
};
|
||||
|
||||
export default config;
|
||||
28
ui/prisma/schema.prisma
Normal file
28
ui/prisma/schema.prisma
Normal file
@@ -0,0 +1,28 @@
|
||||
generator client {
|
||||
provider = "prisma-client-js"
|
||||
}
|
||||
|
||||
datasource db {
|
||||
provider = "sqlite"
|
||||
url = "file:../../aitk_db.db"
|
||||
}
|
||||
|
||||
model Settings {
|
||||
id Int @id @default(autoincrement())
|
||||
key String @unique
|
||||
value String
|
||||
}
|
||||
|
||||
model Job {
|
||||
id String @id @default(uuid())
|
||||
name String @unique
|
||||
gpu_ids String
|
||||
job_config String // JSON string
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
status String @default("stopped")
|
||||
stop Boolean @default(false)
|
||||
step Int @default(0)
|
||||
info String @default("")
|
||||
speed_string String @default("")
|
||||
}
|
||||
1
ui/public/file.svg
Normal file
1
ui/public/file.svg
Normal file
@@ -0,0 +1 @@
|
||||
<svg fill="none" viewBox="0 0 16 16" xmlns="http://www.w3.org/2000/svg"><path d="M14.5 13.5V5.41a1 1 0 0 0-.3-.7L9.8.29A1 1 0 0 0 9.08 0H1.5v13.5A2.5 2.5 0 0 0 4 16h8a2.5 2.5 0 0 0 2.5-2.5m-1.5 0v-7H8v-5H3v12a1 1 0 0 0 1 1h8a1 1 0 0 0 1-1M9.5 5V2.12L12.38 5zM5.13 5h-.62v1.25h2.12V5zm-.62 3h7.12v1.25H4.5zm.62 3h-.62v1.25h7.12V11z" clip-rule="evenodd" fill="#666" fill-rule="evenodd"/></svg>
|
||||
|
After Width: | Height: | Size: 391 B |
1
ui/public/globe.svg
Normal file
1
ui/public/globe.svg
Normal file
@@ -0,0 +1 @@
|
||||
<svg fill="none" xmlns="http://www.w3.org/2000/svg" viewBox="0 0 16 16"><g clip-path="url(#a)"><path fill-rule="evenodd" clip-rule="evenodd" d="M10.27 14.1a6.5 6.5 0 0 0 3.67-3.45q-1.24.21-2.7.34-.31 1.83-.97 3.1M8 16A8 8 0 1 0 8 0a8 8 0 0 0 0 16m.48-1.52a7 7 0 0 1-.96 0H7.5a4 4 0 0 1-.84-1.32q-.38-.89-.63-2.08a40 40 0 0 0 3.92 0q-.25 1.2-.63 2.08a4 4 0 0 1-.84 1.31zm2.94-4.76q1.66-.15 2.95-.43a7 7 0 0 0 0-2.58q-1.3-.27-2.95-.43a18 18 0 0 1 0 3.44m-1.27-3.54a17 17 0 0 1 0 3.64 39 39 0 0 1-4.3 0 17 17 0 0 1 0-3.64 39 39 0 0 1 4.3 0m1.1-1.17q1.45.13 2.69.34a6.5 6.5 0 0 0-3.67-3.44q.65 1.26.98 3.1M8.48 1.5l.01.02q.41.37.84 1.31.38.89.63 2.08a40 40 0 0 0-3.92 0q.25-1.2.63-2.08a4 4 0 0 1 .85-1.32 7 7 0 0 1 .96 0m-2.75.4a6.5 6.5 0 0 0-3.67 3.44 29 29 0 0 1 2.7-.34q.31-1.83.97-3.1M4.58 6.28q-1.66.16-2.95.43a7 7 0 0 0 0 2.58q1.3.27 2.95.43a18 18 0 0 1 0-3.44m.17 4.71q-1.45-.12-2.69-.34a6.5 6.5 0 0 0 3.67 3.44q-.65-1.27-.98-3.1" fill="#666"/></g><defs><clipPath id="a"><path fill="#fff" d="M0 0h16v16H0z"/></clipPath></defs></svg>
|
||||
|
After Width: | Height: | Size: 1.0 KiB |
1
ui/public/next.svg
Normal file
1
ui/public/next.svg
Normal file
@@ -0,0 +1 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" fill="none" viewBox="0 0 394 80"><path fill="#000" d="M262 0h68.5v12.7h-27.2v66.6h-13.6V12.7H262V0ZM149 0v12.7H94v20.4h44.3v12.6H94v21h55v12.6H80.5V0h68.7zm34.3 0h-17.8l63.8 79.4h17.9l-32-39.7 32-39.6h-17.9l-23 28.6-23-28.6zm18.3 56.7-9-11-27.1 33.7h17.8l18.3-22.7z"/><path fill="#000" d="M81 79.3 17 0H0v79.3h13.6V17l50.2 62.3H81Zm252.6-.4c-1 0-1.8-.4-2.5-1s-1.1-1.6-1.1-2.6.3-1.8 1-2.5 1.6-1 2.6-1 1.8.3 2.5 1a3.4 3.4 0 0 1 .6 4.3 3.7 3.7 0 0 1-3 1.8zm23.2-33.5h6v23.3c0 2.1-.4 4-1.3 5.5a9.1 9.1 0 0 1-3.8 3.5c-1.6.8-3.5 1.3-5.7 1.3-2 0-3.7-.4-5.3-1s-2.8-1.8-3.7-3.2c-.9-1.3-1.4-3-1.4-5h6c.1.8.3 1.6.7 2.2s1 1.2 1.6 1.5c.7.4 1.5.5 2.4.5 1 0 1.8-.2 2.4-.6a4 4 0 0 0 1.6-1.8c.3-.8.5-1.8.5-3V45.5zm30.9 9.1a4.4 4.4 0 0 0-2-3.3 7.5 7.5 0 0 0-4.3-1.1c-1.3 0-2.4.2-3.3.5-.9.4-1.6 1-2 1.6a3.5 3.5 0 0 0-.3 4c.3.5.7.9 1.3 1.2l1.8 1 2 .5 3.2.8c1.3.3 2.5.7 3.7 1.2a13 13 0 0 1 3.2 1.8 8.1 8.1 0 0 1 3 6.5c0 2-.5 3.7-1.5 5.1a10 10 0 0 1-4.4 3.5c-1.8.8-4.1 1.2-6.8 1.2-2.6 0-4.9-.4-6.8-1.2-2-.8-3.4-2-4.5-3.5a10 10 0 0 1-1.7-5.6h6a5 5 0 0 0 3.5 4.6c1 .4 2.2.6 3.4.6 1.3 0 2.5-.2 3.5-.6 1-.4 1.8-1 2.4-1.7a4 4 0 0 0 .8-2.4c0-.9-.2-1.6-.7-2.2a11 11 0 0 0-2.1-1.4l-3.2-1-3.8-1c-2.8-.7-5-1.7-6.6-3.2a7.2 7.2 0 0 1-2.4-5.7 8 8 0 0 1 1.7-5 10 10 0 0 1 4.3-3.5c2-.8 4-1.2 6.4-1.2 2.3 0 4.4.4 6.2 1.2 1.8.8 3.2 2 4.3 3.4 1 1.4 1.5 3 1.5 5h-5.8z"/></svg>
|
||||
|
After Width: | Height: | Size: 1.3 KiB |
BIN
ui/public/ostris_logo.png
Normal file
BIN
ui/public/ostris_logo.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 22 KiB |
1
ui/public/vercel.svg
Normal file
1
ui/public/vercel.svg
Normal file
@@ -0,0 +1 @@
|
||||
<svg fill="none" xmlns="http://www.w3.org/2000/svg" viewBox="0 0 1155 1000"><path d="m577.3 0 577.4 1000H0z" fill="#fff"/></svg>
|
||||
|
After Width: | Height: | Size: 128 B |
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user