93 Commits

Author SHA1 Message Date
Jaret Burkett
ed1deb71c4 Added examples for training lumina2 2025-02-08 16:13:18 -07:00
Jaret Burkett
4de6a825fa Update lumina requirements 2025-02-08 15:16:35 -07:00
Jaret Burkett
9a7266275d Wokr on lumina2 2025-02-08 14:52:39 -07:00
Jaret Burkett
d138f07365 Imitial lumina3 support 2025-02-08 10:59:53 -07:00
Jaret Burkett
c6d8eedb94 Added ability to use consistent noise for each image in a dataset by hashing the path and using that as a seed. 2025-02-08 07:13:48 -07:00
Jaret Burkett
af5e760be1 Merge pull request #249 from ostris/accelerate-multi-gpu
Multi gpu support. Other goodies
2025-02-08 07:11:10 -07:00
Jaret Burkett
ff3d54bb5b Make the mean of the mask multiplier be 1.0 for a more balanced loss. 2025-02-06 10:57:06 +00:00
Jaret Burkett
0e75724b4d Lock version of diffusers 2025-02-05 18:14:59 +00:00
Jaret Burkett
376bb1bf6f Lock torch version due to breaking changes 2025-02-04 23:00:57 +00:00
Jaret Burkett
216ab164ce Experimental features and bug fixes 2025-02-04 13:36:34 -07:00
Jaret Burkett
e6180d1e1d Bug fixes 2025-01-31 13:23:01 -07:00
Jaret Burkett
15a57bc89f Add new version of DFE. Kitchen sink 2025-01-31 11:42:27 -07:00
Jaret Burkett
e5355bf8d5 Added train catch to the blank network 2025-01-30 16:15:45 +00:00
Jaret Burkett
34a1c6947a Added flux_shift as timestep type 2025-01-27 07:35:00 -07:00
Jaret Burkett
2141c6e06c Merge remote-tracking branch 'origin/main' into accelerate-multi-gpu 2025-01-26 11:19:34 -07:00
Jaret Burkett
1188cf1e8a Adjust flux sample sampler to handle some new breaking changes in diffusers. 2025-01-26 18:09:21 +00:00
Jaret Burkett
5e663746b8 Working multi gpu training. Still need a lot of tweaks and testing. 2025-01-25 16:46:20 -07:00
Jaret Burkett
441474e81f Added a flag to lora extraction script to do a full transformer extraction. 2025-01-24 09:34:13 -07:00
Jaret Burkett
a6a690f796 Update full fine tune example to only train transformer blocks. 2025-01-24 09:28:34 -07:00
Jaret Burkett
6191f19e55 Added script to convert diffusers model to ComfyUI variant 2025-01-23 21:30:23 -07:00
Jaret Burkett
bbfba0c188 Added v2 of dfp 2025-01-22 16:32:13 -07:00
Jaret Burkett
e1549ad54d Update dfe model arch 2025-01-22 10:37:23 -07:00
Jaret Burkett
04abe57c76 Added weighing to DFE 2025-01-22 08:50:57 -07:00
Jaret Burkett
89dd041b97 Added ability to pair samples with a closer noise with optimal_noise_pairing_samples 2025-01-21 18:30:10 -07:00
Jaret Burkett
29122b1a54 Added code to handle diffusion feature extraction loss 2025-01-21 14:21:34 -07:00
Jaret Burkett
6a8e3d8610 Added a config file for full finetuning flex. Added a lora extraction script for flex 2025-01-20 10:09:01 -07:00
Jaret Burkett
4c8a9e1b88 Added example config to train Flex 2025-01-18 18:03:20 -07:00
Jaret Burkett
fadb2f3a76 Allow quantizing the te independently on flux. added lognorm_blend timestep schedule 2025-01-18 18:02:31 -07:00
Jaret Burkett
4723f23c0d Added ability to split up flux across gpus (experimental). Changed the way timestep scheduling works to prep for more specific schedules. 2024-12-31 07:06:55 -07:00
Jaret Burkett
8ef07a9c36 Added training for an experimental decoratgor embedding. Allow for turning off guidance embedding on flux (for unreleased model). Various bug fixes and modifications 2024-12-15 08:59:27 -07:00
Jaret Burkett
92ce93140e Adjustments to defaults for automagic 2024-11-29 10:28:06 -07:00
Jaret Burkett
f213996aa5 Fixed saving and displaying for automagic 2024-11-29 08:00:22 -07:00
Jaret Burkett
cbe31eaf0a Initial work on a auto adjusting optimizer 2024-11-29 04:48:58 -07:00
Jaret Burkett
67c2e44edb Added support for training flux redux adapters 2024-11-21 20:01:52 -07:00
Jaret Burkett
96d418bb95 Added support for full finetuning flux with randomized param activation. Examples coming soon 2024-11-21 13:05:32 -07:00
Jaret Burkett
894374b2e9 Various bug fixes and optimizations for quantized training. Added untested custom adam8bit optimizer. Did some work on LoRM (dont use) 2024-11-20 09:16:55 -07:00
Jaret Burkett
6509ba4484 Fix seed generation to make it deterministic so it is consistant from gpu to gpu 2024-11-15 12:11:13 -07:00
Jaret Burkett
025ee3dd3d Added ability for adafactor to fully fine tune quantized model. 2024-10-30 16:38:07 -06:00
Jaret Burkett
58f9d01c2b Added adafactor implementation that handles stochastic rounding of update and accumulation 2024-10-30 05:25:57 -06:00
Jaret Burkett
e72b59a8e9 Added experimental 8bit version of prodigy with stochastic rounding and stochastic gradient accumulation. Still testing. 2024-10-29 14:28:28 -06:00
Jaret Burkett
4aa19b5c1d Only quantize flux T5 is also quantizing model. Load TE from original name and path if fine tuning. 2024-10-29 14:25:31 -06:00
Jaret Burkett
4747716867 Fixed issue with adapters not providing gradients with new grad activator 2024-10-29 14:22:10 -06:00
Jaret Burkett
22cd40d7b9 Improvements for full tuning flux. Added debugging launch config for vscode 2024-10-29 04:54:08 -06:00
Jaret Burkett
3400882a80 Added preliminary support for SD3.5-large lora training 2024-10-22 12:21:36 -06:00
Jaret Burkett
9f94c7b61e Added experimental param multiplier to the ema module 2024-10-22 09:25:52 -06:00
Jaret Burkett
bedb8197a2 Fixed issue with sizes for some images being loaded sideways resulting in squished images. 2024-10-20 11:51:29 -06:00
Jaret Burkett
e3ebd73610 Add a projection layer on vision direct when doing image embeds 2024-10-20 10:48:23 -06:00
Jaret Burkett
dd931757cd Merge branch 'main' of github.com:ostris/ai-toolkit 2024-10-20 07:04:29 -06:00
Jaret Burkett
0640cdf569 Handle errors in loading size database 2024-10-20 07:04:19 -06:00
Jaret Burkett
0b048d0dde Locked version of quanto as it breaks in later versions 2024-10-16 22:41:04 +00:00
Jaret Burkett
473d455f44 Process empty clip image if there is not one for reg images when training a custom adapter 2024-10-15 08:28:04 -06:00
Jaret Burkett
ce759ebd8c Normalize the image embeddings on vd adapter forward 2024-10-12 15:09:48 +00:00
Jaret Burkett
628a7923a3 Remove norm on image embeds on custom adapter 2024-10-12 00:43:18 +00:00
Jaret Burkett
3922981996 Added some additional experimental things to the vision direct encoder 2024-10-10 19:42:26 +00:00
Jaret Burkett
ab22674980 Allow for a default caption file in the folder. Minor bug fixes. 2024-10-10 07:31:33 -06:00
Jaret Burkett
9452929300 Apply a mask to the embeds for SD if using T5 encoder 2024-10-04 10:55:20 -06:00
Jaret Burkett
a800c9d19e Add a method to have an inference only lora 2024-10-04 10:06:53 -06:00
Jaret Burkett
28e6f00790 Fixed bug in returning clip image embed to actually return it 2024-10-03 10:49:09 -06:00
Jaret Burkett
67e0aca750 Added ability to load clip pairs randomly from folder. Other small bug fixes 2024-10-03 10:03:49 -06:00
Jaret Burkett
f05224970f Added Vision Languate Adapter usage for pixtral vd adapter 2024-09-29 19:39:56 -06:00
Jaret Burkett
b4f64de4c2 Quick patch to scope xformer imports until a better solution 2024-09-28 15:36:42 -06:00
Jaret Burkett
2e5f6668dc Add xformers ad a dependency 2024-09-28 15:30:14 -06:00
Jaret Burkett
e4c82803e1 Handle random resizing for pixtral input on direct vision adapter 2024-09-28 14:53:38 -06:00
Jaret Burkett
69aa92bce5 Added support for AdEMAMix8bit 2024-09-28 14:33:51 -06:00
Jaret Burkett
a508caad1d Change pixtral to crop based on number of pixels instead of largest dimension 2024-09-28 13:05:26 -06:00
Jaret Burkett
58537fc92b Added initial direct vision pixtral support 2024-09-28 10:47:51 -06:00
Jaret Burkett
86b5938cf3 Fixed the webp bug finally. 2024-09-25 13:56:00 -06:00
Jaret Burkett
6b4034122f REmove layers from direct vision resampler 2024-09-24 15:08:29 -06:00
Jaret Burkett
10817696fb Fixed issue where direct vision was not passing additional params from resampler when it is added 2024-09-24 10:34:11 -06:00
Jaret Burkett
037ce11740 Always return vision encoder in state dict 2024-09-24 07:43:17 -06:00
Jaret Burkett
04424fe2d6 Added config setting to set the timestep type 2024-09-24 06:53:59 -06:00
Jaret Burkett
40a8ff5731 Load local hugging face packages for assistant adapter 2024-09-23 10:37:12 -06:00
Jaret Burkett
2776221497 Added option to cache empty prompt or trigger and unload text encoders while training 2024-09-21 20:54:09 -06:00
Jaret Burkett
f85ad452c6 Added initial support for pixtral vision as a vision encoder. 2024-09-21 15:21:14 -06:00
Jaret Burkett
dd889086f4 Updates to the docker file for jupyterlab 2024-09-21 12:07:07 -06:00
apolinário
bc693488eb fix diffusers codebase (#183) 2024-09-21 11:50:29 -06:00
Jaret Burkett
d97c55cd96 Updated requirements to lock version of albucore, which had breaking changes. 2024-09-21 11:19:13 -06:00
Plat
79b4e04b80 Feat: Wandb logging (#95)
* wandb logging

* fix: start logging before train loop

* chore: add wandb dir to gitignore

* fix: wrap wandb functions

* fix: forget to send last samples

* chore: use valid type

* chore: use None when not type-checking

* chore: resolved complicated logic

* fix: follow log_every

---------

Co-authored-by: Plat <github@p1at.dev>
Co-authored-by: Jaret Burkett <jaretburkett@gmail.com>
2024-09-19 20:01:01 -06:00
Jaret Burkett
951e223481 Added support to disable single transformers in vision direct adapter 2024-09-11 08:54:51 -06:00
Jaret Burkett
fc34a69bec Ignore guidance embed when full tuning flux. adjust block scaler to decat to 1.0. Add MLP resampler for reducing vision adapter tokens 2024-09-09 16:24:46 -06:00
Jaret Burkett
279ee65177 Remove block scaler 2024-09-06 08:28:17 -06:00
Jaret Burkett
3a1f464132 Added support for training vision direct weight adapters 2024-09-05 10:11:44 -06:00
Jaret Burkett
5c8fcc8a4e Fix bug with zeroing out gradients when accumulating 2024-09-03 08:29:15 -06:00
Jaret Burkett
121a760c19 Added proper grad accumulation 2024-09-03 07:24:18 -06:00
Jaret Burkett
e5fadddd45 Added ability to do prompt attn masking for flux 2024-09-02 17:29:36 -06:00
Jaret Burkett
d44d4eb61a Added a new experimental linear weighing technique 2024-09-02 09:22:13 -06:00
Jaret Burkett
7d9ab22405 Rework ip adapter and vision direct adapters to apply to the single transformer blocks even though they are not cross attn. 2024-09-01 10:40:42 -06:00
Jaret Burkett
7ed8c51f20 Readme cleanup 2024-09-01 07:06:09 -06:00
Jaret Burkett
6df33156f0 Add information about specific weight targeting in the README 2024-09-01 06:59:47 -06:00
Jaret Burkett
40f5c59da0 Fixes for training ilora on flux 2024-08-31 16:55:26 -06:00
Jaret Burkett
3e71a99df0 Check for contains only against clean name for lora, not the adjusted one 2024-08-31 07:44:13 -06:00
apolinário
562405923f Update README.md for push_to_hub (#143)
Add diffusers examples and clarify how to use the model locally
2024-08-30 16:34:28 -06:00
apolinário
f84bd6d7a6 Add Gradio UI for ai-toolkit (#141)
* Add Gradio UI for FLUX.1

* small text changes

* no flash-attn? no problem!

* bye flash-attn!

* fixes for windows

---------

Co-authored-by: multimodalart <joaopaulo.passos+multimodal@gmail.com>
2024-08-30 06:29:51 -06:00
56 changed files with 7394 additions and 892 deletions

6
.gitignore vendored
View File

@@ -173,4 +173,8 @@ cython_debug/
!/output/.gitkeep
/extensions/*
!/extensions/example
/temp
/temp
/wandb
.vscode/settings.json
.DS_Store
._.DS_Store

28
.vscode/launch.json vendored Normal file
View 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
},
]
}

View File

@@ -117,6 +117,20 @@ Please do not open a bug report unless it is a bug in the code. You are welcome
and ask for help there. However, please refrain from PMing me directly with general question or support. Ask in the discord
and I will answer when I can.
## Gradio UI
To get started training locally with a with a custom UI, once you followed the steps above and `ai-toolkit` is installed:
```bash
cd ai-toolkit #in case you are not yet in the ai-toolkit folder
huggingface-cli login #provide a `write` token to publish your LoRA at the end
python flux_train_ui.py
```
You will instantiate a UI that will let you upload your images, caption them, train and publish your LoRA
![image](assets/lora_ease_ui.png)
## Training in RunPod
Example RunPod template: **runpod/pytorch:2.2.0-py3.10-cuda12.1.1-devel-ubuntu22.04**
> You need a minimum of 24GB VRAM, pick a GPU by your preference.
@@ -222,6 +236,54 @@ replaced.
Images are never upscaled but they are downscaled and placed in buckets for batching. **You do not need to crop/resize your images**.
The loader will automatically resize them and can handle varying aspect ratios.
## Training Specific Layers
To train specific layers with LoRA, you can use the `only_if_contains` network kwargs. For instance, if you want to train only the 2 layers
used by The Last Ben, [mentioned in this post](https://x.com/__TheBen/status/1829554120270987740), you can adjust your
network kwargs like so:
```yaml
network:
type: "lora"
linear: 128
linear_alpha: 128
network_kwargs:
only_if_contains:
- "transformer.single_transformer_blocks.7.proj_out"
- "transformer.single_transformer_blocks.20.proj_out"
```
The naming conventions of the layers are in diffusers format, so checking the state dict of a model will reveal
the suffix of the name of the layers you want to train. You can also use this method to only train specific groups of weights.
For instance to only train the `single_transformer` for FLUX.1, you can use the following:
```yaml
network:
type: "lora"
linear: 128
linear_alpha: 128
network_kwargs:
only_if_contains:
- "transformer.single_transformer_blocks."
```
You can also exclude layers by their names by using `ignore_if_contains` network kwarg. So to exclude all the single transformer blocks,
```yaml
network:
type: "lora"
linear: 128
linear_alpha: 128
network_kwargs:
ignore_if_contains:
- "transformer.single_transformer_blocks."
```
`ignore_if_contains` takes priority over `only_if_contains`. So if a weight is covered by both,
if will be ignored.
---
## EVERYTHING BELOW THIS LINE IS OUTDATED

BIN
assets/lora_ease_ui.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 340 KiB

View 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'

View 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'

View 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'

View 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'

View 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'

View File

@@ -1,4 +1,5 @@
FROM runpod/base:0.6.2-cuda12.1.0
FROM runpod/base:0.6.2-cuda12.2.0
LABEL authors="jaret"
# Install dependencies
@@ -17,5 +18,14 @@ RUN python -m pip install -r requirements.txt
RUN apt-get install -y tmux nvtop htop
RUN pip install jupyterlab
# mask workspace
RUN mkdir /workspace
# symlink app to workspace
RUN ln -s /app/ai-toolkit /workspace/ai-toolkit
WORKDIR /
CMD ["/start.sh"]

View File

@@ -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,7 @@ 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
def flush():
@@ -58,23 +60,26 @@ 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.dfe: Optional[DiffusionFeatureExtractor] = None
def before_model_load(self):
@@ -113,6 +118,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 +160,33 @@ 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)
# 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 +292,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
@@ -340,7 +375,46 @@ class SDTrainer(BaseSDTrainProcess):
target = (noise - batch.latents).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_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
@@ -390,9 +464,12 @@ class SDTrainer(BaseSDTrainProcess):
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
@@ -417,7 +494,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])
@@ -457,7 +534,7 @@ class SDTrainer(BaseSDTrainProcess):
loss = loss + norm_std_loss
return loss
return loss + additional_loss
def preprocess_batch(self, batch: 'DataLoaderBatchDTO'):
return batch
@@ -486,7 +563,6 @@ class SDTrainer(BaseSDTrainProcess):
noise=noise,
sd=self.sd,
unconditional_embeds=unconditional_embeds,
scaler=self.scaler,
**kwargs
)
@@ -601,7 +677,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 +832,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 +870,8 @@ class SDTrainer(BaseSDTrainProcess):
was_adapter_active = self.adapter.is_active
self.adapter.is_active = False
if self.train_config.unload_text_encoder:
raise ValueError("Prior predictions currently do not support unloading text encoder")
# 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 +916,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 +980,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)
@@ -1016,6 +1096,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 +1152,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 +1247,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 +1260,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
@@ -1235,12 +1339,23 @@ class SDTrainer(BaseSDTrainProcess):
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 +1395,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 +1457,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:
@@ -1421,7 +1537,6 @@ 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:
quad_count = random.randint(1, 4)
@@ -1447,10 +1562,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 +1592,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:
@@ -1520,10 +1636,9 @@ class SDTrainer(BaseSDTrainProcess):
)
# 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 +1651,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()

414
flux_train_ui.py Normal file
View 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)

View File

@@ -34,6 +34,7 @@ from toolkit.lora_special import LoRASpecialNetwork
from toolkit.lorm import convert_diffusers_unet_to_lorm, count_parameters, print_lorm_extract_details, \
lorm_ignore_if_contains, lorm_parameter_threshold, LORM_TARGET_REPLACE_MODULE
from toolkit.lycoris_special import LycorisSpecialNetwork
from toolkit.models.decorator import Decorator
from toolkit.network_mixins import Network
from toolkit.optimizer import get_optimizer
from toolkit.paths import CONFIG_ROOT
@@ -55,9 +56,17 @@ import gc
from tqdm import tqdm
from toolkit.config_modules import SaveConfig, LogingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \
GenerateImageConfig, EmbeddingConfig, DatasetConfig, preprocess_dataset_raw_config, AdapterConfig, GuidanceConfig
from toolkit.config_modules import SaveConfig, LoggingConfig, SampleConfig, NetworkConfig, TrainConfig, ModelConfig, \
GenerateImageConfig, EmbeddingConfig, DatasetConfig, preprocess_dataset_raw_config, AdapterConfig, GuidanceConfig, validate_configs, \
DecoratorConfig
from toolkit.logging import create_logger
from diffusers import FluxTransformer2DModel
from toolkit.accelerator import get_accelerator
from toolkit.print import print_acc
from accelerate import Accelerator
import transformers
import diffusers
import hashlib
def flush():
torch.cuda.empty_cache()
@@ -68,6 +77,14 @@ class BaseSDTrainProcess(BaseTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict, custom_pipeline=None):
super().__init__(process_id, job, config)
self.accelerator: Accelerator = get_accelerator()
if self.accelerator.is_local_main_process:
transformers.utils.logging.set_verbosity_warning()
diffusers.utils.logging.set_verbosity_error()
else:
transformers.utils.logging.set_verbosity_error()
diffusers.utils.logging.set_verbosity_error()
self.sd: StableDiffusion
self.embedding: Union[Embedding, None] = None
@@ -79,8 +96,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.grad_accumulation_step = 1
# if true, then we do not do an optimizer step. We are accumulating gradients
self.is_grad_accumulation_step = False
self.device = self.get_conf('device', self.job.device)
self.device_torch = torch.device(self.device)
self.device = str(self.accelerator.device)
self.device_torch = self.accelerator.device
network_config = self.get_conf('network', None)
if network_config is not None:
self.network_config = NetworkConfig(**network_config)
@@ -88,6 +105,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.network_config = None
self.train_config = TrainConfig(**self.get_conf('train', {}))
model_config = self.get_conf('model', {})
self.modules_being_trained: List[torch.nn.Module] = []
# update modelconfig dtype to match train
model_config['dtype'] = self.train_config.dtype
@@ -102,7 +120,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
else:
self.has_first_sample_requested = False
self.first_sample_config = self.sample_config
self.logging_config = LogingConfig(**self.get_conf('logging', {}))
self.logging_config = LoggingConfig(**self.get_conf('logging', {}))
self.logger = create_logger(self.logging_config, config)
self.optimizer: torch.optim.Optimizer = None
self.lr_scheduler = None
self.data_loader: Union[DataLoader, None] = None
@@ -141,6 +160,13 @@ class BaseSDTrainProcess(BaseTrainProcess):
embedding_raw = self.get_conf('embedding', None)
if embedding_raw is not None:
self.embed_config = EmbeddingConfig(**embedding_raw)
self.decorator_config: DecoratorConfig = None
decorator_raw = self.get_conf('decorator', None)
if decorator_raw is not None:
if not self.model_config.is_flux:
raise ValueError("Decorators are only supported for Flux models currently")
self.decorator_config = DecoratorConfig(**decorator_raw)
# t2i adapter
self.adapter_config = None
@@ -155,6 +181,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.network: Union[Network, None] = None
self.adapter: Union[T2IAdapter, IPAdapter, ClipVisionAdapter, ReferenceAdapter, CustomAdapter, ControlNetModel, None] = None
self.embedding: Union[Embedding, None] = None
self.decorator: Union[Decorator, None] = None
is_training_adapter = self.adapter_config is not None and self.adapter_config.train
@@ -172,12 +199,29 @@ class BaseSDTrainProcess(BaseTrainProcess):
train_lora=self.network_config is not None,
train_adapter=is_training_adapter,
train_embedding=self.embed_config is not None,
train_decorator=self.decorator_config is not None,
train_refiner=self.train_config.train_refiner,
unload_text_encoder=self.train_config.unload_text_encoder,
require_grads=False # we ensure them later
)
self.get_params_device_state_preset = get_train_sd_device_state_preset(
device=self.device_torch,
train_unet=self.train_config.train_unet,
train_text_encoder=self.train_config.train_text_encoder,
cached_latents=self.is_latents_cached,
train_lora=self.network_config is not None,
train_adapter=is_training_adapter,
train_embedding=self.embed_config is not None,
train_decorator=self.decorator_config is not None,
train_refiner=self.train_config.train_refiner,
unload_text_encoder=self.train_config.unload_text_encoder,
require_grads=True # We check for grads when getting params
)
# fine_tuning here is for training actual SD network, not LoRA, embeddings, etc. it is (Dreambooth, etc)
self.is_fine_tuning = True
if self.network_config is not None or is_training_adapter or self.embed_config is not None:
if self.network_config is not None or is_training_adapter or self.embed_config is not None or self.decorator_config is not None:
self.is_fine_tuning = False
self.named_lora = False
@@ -185,12 +229,16 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.named_lora = True
self.snr_gos: Union[LearnableSNRGamma, None] = None
self.ema: ExponentialMovingAverage = None
validate_configs(self.train_config, self.model_config, self.save_config)
def post_process_generate_image_config_list(self, generate_image_config_list: List[GenerateImageConfig]):
# override in subclass
return generate_image_config_list
def sample(self, step=None, is_first=False):
if not self.accelerator.is_main_process:
return
flush()
sample_folder = os.path.join(self.save_root, 'samples')
gen_img_config_list = []
@@ -258,6 +306,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
adapter_conditioning_scale=sample_config.adapter_conditioning_scale,
refiner_start_at=sample_config.refiner_start_at,
extra_values=sample_config.extra_values,
logger=self.logger,
**extra_args
))
@@ -284,6 +333,10 @@ class BaseSDTrainProcess(BaseTrainProcess):
elif self.model_config.is_xl:
o_dict['ss_base_model_version'] = 'sdxl_1.0'
elif self.model_config.is_flux:
o_dict['ss_base_model_version'] = 'flux.1'
elif self.model_config.is_lumina2:
o_dict['ss_base_model_version'] = 'lumina2'
else:
o_dict['ss_base_model_version'] = 'sd_1.5'
@@ -312,6 +365,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
return info
def clean_up_saves(self):
if not self.accelerator.is_main_process:
return
# remove old saves
# get latest saved step
latest_item = None
@@ -368,7 +423,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
items_to_remove = list(dict.fromkeys(items_to_remove))
for item in items_to_remove:
self.print(f"Removing old save: {item}")
print_acc(f"Removing old save: {item}")
if os.path.isdir(item):
shutil.rmtree(item)
else:
@@ -386,6 +441,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
pass
def save(self, step=None):
if not self.accelerator.is_main_process:
return
flush()
if self.ema is not None:
# always save params as ema
@@ -448,6 +505,19 @@ class BaseSDTrainProcess(BaseTrainProcess):
# replace extension
emb_file_path = os.path.splitext(emb_file_path)[0] + ".pt"
self.embedding.save(emb_file_path)
if self.decorator is not None:
dec_filename = f'{self.job.name}{step_num}.safetensors'
dec_file_path = os.path.join(self.save_root, dec_filename)
decorator_state_dict = self.decorator.state_dict()
for key, value in decorator_state_dict.items():
if isinstance(value, torch.Tensor):
decorator_state_dict[key] = value.clone().to('cpu', dtype=get_torch_dtype(self.save_config.dtype))
save_file(
decorator_state_dict,
dec_file_path,
metadata=save_meta,
)
if self.adapter is not None and self.adapter_config.train:
adapter_name = self.job.name
@@ -493,12 +563,17 @@ class BaseSDTrainProcess(BaseTrainProcess):
# move it back
self.adapter = self.adapter.to(orig_device, dtype=orig_dtype)
else:
direct_save = False
if self.adapter_config.train_only_image_encoder:
direct_save = True
if self.adapter_config.type == 'redux':
direct_save = True
save_ip_adapter_from_diffusers(
state_dict,
output_file=file_path,
meta=save_meta,
dtype=get_torch_dtype(self.save_config.dtype),
direct_save=self.adapter_config.train_only_image_encoder
direct_save=direct_save
)
else:
if self.save_config.save_format == "diffusers":
@@ -541,12 +616,13 @@ class BaseSDTrainProcess(BaseTrainProcess):
try:
filename = f'optimizer.pt'
file_path = os.path.join(self.save_root, filename)
torch.save(self.optimizer.state_dict(), file_path)
state_dict = self.optimizer.state_dict()
torch.save(state_dict, file_path)
except Exception as e:
print(e)
print("Could not save optimizer")
print_acc(e)
print_acc("Could not save optimizer")
self.print(f"Saved to {file_path}")
print_acc(f"Saved to {file_path}")
self.clean_up_saves()
self.post_save_hook(file_path)
@@ -568,13 +644,58 @@ class BaseSDTrainProcess(BaseTrainProcess):
return params
def hook_before_train_loop(self):
pass
if self.accelerator.is_main_process:
self.logger.start()
self.prepare_accelerator()
def prepare_accelerator(self):
# set some config
self.accelerator.even_batches=False
# # prepare all the models stuff for accelerator (hopefully we dont miss any)
self.sd.vae = self.accelerator.prepare(self.sd.vae)
if self.sd.unet is not None:
self.sd.unet_unwrapped = self.sd.unet
self.sd.unet = self.accelerator.prepare(self.sd.unet)
# todo always tdo it?
self.modules_being_trained.append(self.sd.unet)
if self.sd.text_encoder is not None and self.train_config.train_text_encoder:
if isinstance(self.sd.text_encoder, list):
self.sd.text_encoder = [self.accelerator.prepare(model) for model in self.sd.text_encoder]
self.modules_being_trained.extend(self.sd.text_encoder)
else:
self.sd.text_encoder = self.accelerator.prepare(self.sd.text_encoder)
self.modules_being_trained.append(self.sd.text_encoder)
if self.sd.refiner_unet is not None and self.train_config.train_refiner:
self.sd.refiner_unet = self.accelerator.prepare(self.sd.refiner_unet)
self.modules_being_trained.append(self.sd.refiner_unet)
# todo, do we need to do the network or will "unet" get it?
if self.sd.network is not None:
self.sd.network = self.accelerator.prepare(self.sd.network)
self.modules_being_trained.append(self.sd.network)
if self.adapter is not None and self.adapter_config.train:
# todo adapters may not be a module. need to check
self.adapter = self.accelerator.prepare(self.adapter)
self.modules_being_trained.append(self.adapter)
# prepare other things
self.optimizer = self.accelerator.prepare(self.optimizer)
if self.lr_scheduler is not None:
self.lr_scheduler = self.accelerator.prepare(self.lr_scheduler)
# self.data_loader = self.accelerator.prepare(self.data_loader)
# if self.data_loader_reg is not None:
# self.data_loader_reg = self.accelerator.prepare(self.data_loader_reg)
def ensure_params_requires_grad(self):
# get param groups
for group in self.optimizer.param_groups:
def ensure_params_requires_grad(self, force=False):
if self.train_config.do_paramiter_swapping and not force:
# the optimizer will handle this if we are not forcing
return
for group in self.params:
for param in group['params']:
param.requires_grad = True
if isinstance(param, torch.nn.Parameter): # Ensure it's a proper parameter
param.requires_grad_(True)
def setup_ema(self):
if self.train_config.ema_config.use_ema:
@@ -585,8 +706,9 @@ class BaseSDTrainProcess(BaseTrainProcess):
params.append(param)
self.ema = ExponentialMovingAverage(
params,
self.train_config.ema_config.ema_decay,
decay=self.train_config.ema_config.ema_decay,
use_feedback=self.train_config.ema_config.use_feedback,
param_multiplier=self.train_config.ema_config.param_multiplier,
)
def before_dataset_load(self):
@@ -637,6 +759,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
return latest_path
def load_training_state_from_metadata(self, path):
if not self.accelerator.is_main_process:
return
meta = None
# if path is folder, then it is diffusers
if os.path.isdir(path):
@@ -653,7 +777,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
if 'epoch' in meta['training_info']:
self.epoch_num = meta['training_info']['epoch']
self.start_step = self.step_num
print(f"Found step {self.step_num} in metadata, starting from there")
print_acc(f"Found step {self.step_num} in metadata, starting from there")
def load_weights(self, path):
if self.network is not None:
@@ -661,7 +785,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.load_training_state_from_metadata(path)
return extra_weights
else:
print("load_weights not implemented for non-network models")
print_acc("load_weights not implemented for non-network models")
return None
def apply_snr(self, seperated_loss, timesteps):
@@ -692,7 +816,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
if 'epoch' in meta['training_info']:
self.epoch_num = meta['training_info']['epoch']
self.start_step = self.step_num
print(f"Found step {self.step_num} in metadata, starting from there")
print_acc(f"Found step {self.step_num} in metadata, starting from there")
# def get_sigmas(self, timesteps, n_dim=4, dtype=torch.float32):
# self.sd.noise_scheduler.set_timesteps(1000, device=self.device_torch)
@@ -723,15 +847,61 @@ class BaseSDTrainProcess(BaseTrainProcess):
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
def get_optimal_noise(self, latents, dtype=torch.float32):
batch_num = latents.shape[0]
chunks = torch.chunk(latents, batch_num, dim=0)
noise_chunks = []
for chunk in chunks:
noise_samples = [torch.randn_like(chunk, device=chunk.device, dtype=dtype) for _ in range(self.train_config.optimal_noise_pairing_samples)]
# find the one most similar to the chunk
lowest_loss = 999999999999
best_noise = None
for noise in noise_samples:
loss = torch.nn.functional.mse_loss(chunk, noise)
if loss < lowest_loss:
lowest_loss = loss
best_noise = noise
noise_chunks.append(best_noise)
noise = torch.cat(noise_chunks, dim=0)
return noise
def get_consistent_noise(self, latents, batch: 'DataLoaderBatchDTO', dtype=torch.float32):
batch_num = latents.shape[0]
chunks = torch.chunk(latents, batch_num, dim=0)
noise_chunks = []
for idx, chunk in enumerate(chunks):
# get seed from path
file_item = batch.file_items[idx]
img_path = file_item.path
# add augmentors
if file_item.flip_x:
img_path += '_fx'
if file_item.flip_y:
img_path += '_fy'
seed = int(hashlib.md5(img_path.encode()).hexdigest(), 16) & 0xffffffff
generator = torch.Generator("cpu").manual_seed(seed)
noise_chunk = torch.randn(chunk.shape, generator=generator).to(chunk.device, dtype=dtype)
noise_chunks.append(noise_chunk)
noise = torch.cat(noise_chunks, dim=0).to(dtype=dtype)
return noise
def get_noise(self, latents, batch_size, dtype=torch.float32):
# get noise
noise = self.sd.get_latent_noise(
height=latents.shape[2],
width=latents.shape[3],
batch_size=batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
def get_noise(self, latents, batch_size, dtype=torch.float32, batch: 'DataLoaderBatchDTO' = None):
if self.train_config.optimal_noise_pairing_samples > 1:
noise = self.get_optimal_noise(latents, dtype=dtype)
elif self.train_config.force_consistent_noise:
if batch is None:
raise ValueError("Batch must be provided for consistent noise")
noise = self.get_consistent_noise(latents, batch, dtype=dtype)
else:
# get noise
noise = self.sd.get_latent_noise(
height=latents.shape[2],
width=latents.shape[3],
batch_size=batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
if self.train_config.random_noise_shift > 0.0:
# get random noise -1 to 1
@@ -910,10 +1080,21 @@ class BaseSDTrainProcess(BaseTrainProcess):
num_train_timesteps, device=self.device_torch, original_inference_steps=num_train_timesteps
)
elif 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
self.sd.noise_scheduler.set_train_timesteps(
num_train_timesteps,
device=self.device_torch,
linear=self.train_config.linear_timesteps
timestep_type=timestep_type,
latents=latents
)
else:
self.sd.noise_scheduler.set_timesteps(
@@ -983,7 +1164,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
timesteps = torch.stack(timesteps, dim=0)
# get noise
noise = self.get_noise(latents, batch_size, dtype=dtype)
noise = self.get_noise(latents, batch_size, dtype=dtype, batch=batch)
# add dynamic noise offset. Dynamic noise is offsetting the noise to the same channelwise mean as the latents
# this will negate any noise offsets
@@ -1157,7 +1338,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.adapter.to(self.device_torch, dtype=dtype)
if latest_save_path is not None and not is_control_net:
# load adapter from path
print(f"Loading adapter from {latest_save_path}")
print_acc(f"Loading adapter from {latest_save_path}")
if is_t2i:
loaded_state_dict = load_t2i_model(
latest_save_path,
@@ -1203,17 +1384,24 @@ class BaseSDTrainProcess(BaseTrainProcess):
latest_save_path = self.get_latest_save_path()
if latest_save_path is not None:
print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
print_acc(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
model_config_to_load.name_or_path = latest_save_path
self.load_training_state_from_metadata(latest_save_path)
# get the noise scheduler
arch = 'sd'
if self.model_config.is_pixart:
arch = 'pixart'
if self.model_config.is_flux:
arch = 'flux'
if self.model_config.is_lumina2:
arch = 'lumina2'
sampler = get_sampler(
self.train_config.noise_scheduler,
{
"prediction_type": "v_prediction" if self.model_config.is_v_pred else "epsilon",
},
'sd' if not self.model_config.is_pixart else 'pixart'
arch=arch,
)
if self.train_config.train_refiner and self.model_config.refiner_name_or_path is not None and self.network_config is None:
@@ -1253,12 +1441,33 @@ class BaseSDTrainProcess(BaseTrainProcess):
torch.backends.cuda.enable_math_sdp(True)
torch.backends.cuda.enable_flash_sdp(True)
torch.backends.cuda.enable_mem_efficient_sdp(True)
# # check if we have sage and is flux
# if self.sd.is_flux:
# # try_to_activate_sage_attn()
# try:
# from sageattention import sageattn
# from toolkit.models.flux_sage_attn import FluxSageAttnProcessor2_0
# model: FluxTransformer2DModel = self.sd.unet
# # enable sage attention on each block
# for block in model.transformer_blocks:
# processor = FluxSageAttnProcessor2_0()
# block.attn.set_processor(processor)
# for block in model.single_transformer_blocks:
# processor = FluxSageAttnProcessor2_0()
# block.attn.set_processor(processor)
# except ImportError:
# print_acc("sage attention is not installed. Using SDP instead")
if self.train_config.gradient_checkpointing:
if self.sd.is_flux:
# if has method enable_gradient_checkpointing
if hasattr(unet, 'enable_gradient_checkpointing'):
unet.enable_gradient_checkpointing()
elif hasattr(unet, 'gradient_checkpointing'):
unet.gradient_checkpointing = True
else:
unet.enable_gradient_checkpointing()
print("Gradient checkpointing not supported on this model")
if isinstance(text_encoder, list):
for te in text_encoder:
if hasattr(te, 'enable_gradient_checkpointing'):
@@ -1350,6 +1559,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
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,
@@ -1426,8 +1636,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
latest_save_path = self.get_latest_save_path(lora_name)
extra_weights = None
if latest_save_path is not None:
self.print(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
self.print(f"Loading from {latest_save_path}")
print_acc(f"#### IMPORTANT RESUMING FROM {latest_save_path} ####")
print_acc(f"Loading from {latest_save_path}")
extra_weights = self.load_weights(latest_save_path)
self.network.multiplier = 1.0
@@ -1448,11 +1658,35 @@ class BaseSDTrainProcess(BaseTrainProcess):
# self.step_num = self.embedding.step
# self.start_step = self.step_num
params.append({
'params': self.embedding.get_trainable_params(),
'params': list(self.embedding.get_trainable_params()),
'lr': self.train_config.embedding_lr
})
flush()
if self.decorator_config is not None:
self.decorator = Decorator(
num_tokens=self.decorator_config.num_tokens,
token_size=4096 # t5xxl hidden size for flux
)
latest_save_path = self.get_latest_save_path()
# load last saved weights
if latest_save_path is not None:
state_dict = load_file(latest_save_path)
self.decorator.load_state_dict(state_dict)
self.load_training_state_from_metadata(latest_save_path)
params.append({
'params': list(self.decorator.parameters()),
'lr': self.train_config.lr
})
# give it to the sd network
self.sd.decorator = self.decorator
self.decorator.to(self.device_torch, dtype=torch.float32)
self.decorator.train()
flush()
if self.adapter_config is not None:
self.setup_adapter()
@@ -1466,7 +1700,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
else:
# set trainable params
params.append({
'params': self.adapter.parameters(),
'params': list(self.adapter.parameters()),
'lr': self.train_config.adapter_lr
})
@@ -1478,7 +1712,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
else: # no network, embedding or adapter
# set the device state preset before getting params
self.sd.set_device_state(self.train_device_state_preset)
self.sd.set_device_state(self.get_params_device_state_preset)
# params = self.get_params()
if len(params) == 0:
@@ -1512,9 +1746,17 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.start_step = self.step_num
optimizer_type = self.train_config.optimizer.lower()
# esure params require grad
self.ensure_params_requires_grad(force=True)
optimizer = get_optimizer(self.params, optimizer_type, learning_rate=self.train_config.lr,
optimizer_params=self.train_config.optimizer_params)
self.optimizer = optimizer
# set it to do paramiter swapping
if self.train_config.do_paramiter_swapping:
# only works for adafactor, but it should have thrown an error prior to this otherwise
self.optimizer.enable_paramiter_swapping(self.train_config.paramiter_swapping_factor)
# check if it exists
optimizer_state_filename = f'optimizer.pt'
@@ -1528,17 +1770,17 @@ class BaseSDTrainProcess(BaseTrainProcess):
previous_lrs.append(group['lr'])
try:
print(f"Loading optimizer state from {optimizer_state_file_path}")
print_acc(f"Loading optimizer state from {optimizer_state_file_path}")
optimizer_state_dict = torch.load(optimizer_state_file_path, weights_only=True)
optimizer.load_state_dict(optimizer_state_dict)
del optimizer_state_dict
flush()
except Exception as e:
print(f"Failed to load optimizer state from {optimizer_state_file_path}")
print(e)
print_acc(f"Failed to load optimizer state from {optimizer_state_file_path}")
print_acc(e)
# update the optimizer LR from the params
print(f"Updating optimizer LR from params")
print_acc(f"Updating optimizer LR from params")
if len(previous_lrs) > 0:
for i, group in enumerate(optimizer.param_groups):
group['lr'] = previous_lrs[i]
@@ -1574,24 +1816,27 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.hook_before_train_loop()
if self.has_first_sample_requested and self.step_num <= 1 and not self.train_config.disable_sampling:
self.print("Generating first sample from first sample config")
print_acc("Generating first sample from first sample config")
self.sample(0, is_first=True)
# sample first
if self.train_config.skip_first_sample or self.train_config.disable_sampling:
self.print("Skipping first sample due to config setting")
print_acc("Skipping first sample due to config setting")
elif self.step_num <= 1 or self.train_config.force_first_sample:
self.print("Generating baseline samples before training")
print_acc("Generating baseline samples before training")
self.sample(self.step_num)
self.progress_bar = ToolkitProgressBar(
total=self.train_config.steps,
desc=self.job.name,
leave=True,
initial=self.step_num,
iterable=range(0, self.train_config.steps),
)
self.progress_bar.pause()
if self.accelerator.is_local_main_process:
self.progress_bar = ToolkitProgressBar(
total=self.train_config.steps,
desc=self.job.name,
leave=True,
initial=self.step_num,
iterable=range(0, self.train_config.steps),
)
self.progress_bar.pause()
else:
self.progress_bar = None
if self.data_loader is not None:
dataloader = self.data_loader
@@ -1616,20 +1861,23 @@ class BaseSDTrainProcess(BaseTrainProcess):
flush()
# self.step_num = 0
# print(f"Compiling Model")
# print_acc(f"Compiling Model")
# torch.compile(self.sd.unet, dynamic=True)
# make sure all params require grad
self.ensure_params_requires_grad()
self.ensure_params_requires_grad(force=True)
###################################################################
# TRAIN LOOP
###################################################################
start_step_num = self.step_num
did_first_flush = False
for step in range(start_step_num, self.train_config.steps):
if self.train_config.do_paramiter_swapping:
self.optimizer.optimizer.swap_paramiters()
self.timer.start('train_loop')
if self.train_config.do_random_cfg:
self.train_config.do_cfg = True
@@ -1639,7 +1887,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
self.is_grad_accumulation_step = True
if self.train_config.free_u:
self.sd.pipeline.enable_freeu(s1=0.9, s2=0.2, b1=1.1, b2=1.2)
self.progress_bar.unpause()
if self.progress_bar is not None:
self.progress_bar.unpause()
with torch.no_grad():
# if is even step and we have a reg dataset, use that
# todo improve this logic to send one of each through if we can buckets and batch size might be an issue
@@ -1648,42 +1897,54 @@ class BaseSDTrainProcess(BaseTrainProcess):
is_sample_step = self.sample_config.sample_every and self.step_num % self.sample_config.sample_every == 0
if self.train_config.disable_sampling:
is_sample_step = False
# don't do a reg step on sample or save steps as we dont want to normalize on those
if step % 2 == 0 and dataloader_reg is not None and not is_save_step and not is_sample_step:
try:
with self.timer('get_batch:reg'):
batch = next(dataloader_iterator_reg)
except StopIteration:
with self.timer('reset_batch:reg'):
# hit the end of an epoch, reset
self.progress_bar.pause()
dataloader_iterator_reg = iter(dataloader_reg)
trigger_dataloader_setup_epoch(dataloader_reg)
with self.timer('get_batch:reg'):
batch = next(dataloader_iterator_reg)
self.progress_bar.unpause()
is_reg_step = True
elif dataloader is not None:
try:
with self.timer('get_batch'):
batch = next(dataloader_iterator)
except StopIteration:
with self.timer('reset_batch'):
# hit the end of an epoch, reset
self.progress_bar.pause()
dataloader_iterator = iter(dataloader)
trigger_dataloader_setup_epoch(dataloader)
self.epoch_num += 1
if self.train_config.gradient_accumulation_steps == -1:
# if we are accumulating for an entire epoch, trigger a step
self.is_grad_accumulation_step = False
self.grad_accumulation_step = 0
with self.timer('get_batch'):
batch = next(dataloader_iterator)
self.progress_bar.unpause()
else:
batch = None
batch_list = []
for b in range(self.train_config.gradient_accumulation):
# keep track to alternate on an accumulation step for reg
batch_step = step
# don't do a reg step on sample or save steps as we dont want to normalize on those
if batch_step % 2 == 0 and dataloader_reg is not None and not is_save_step and not is_sample_step:
try:
with self.timer('get_batch:reg'):
batch = next(dataloader_iterator_reg)
except StopIteration:
with self.timer('reset_batch:reg'):
# hit the end of an epoch, reset
if self.progress_bar is not None:
self.progress_bar.pause()
dataloader_iterator_reg = iter(dataloader_reg)
trigger_dataloader_setup_epoch(dataloader_reg)
with self.timer('get_batch:reg'):
batch = next(dataloader_iterator_reg)
if self.progress_bar is not None:
self.progress_bar.unpause()
is_reg_step = True
elif dataloader is not None:
try:
with self.timer('get_batch'):
batch = next(dataloader_iterator)
except StopIteration:
with self.timer('reset_batch'):
# hit the end of an epoch, reset
if self.progress_bar is not None:
self.progress_bar.pause()
dataloader_iterator = iter(dataloader)
trigger_dataloader_setup_epoch(dataloader)
self.epoch_num += 1
if self.train_config.gradient_accumulation_steps == -1:
# if we are accumulating for an entire epoch, trigger a step
self.is_grad_accumulation_step = False
self.grad_accumulation_step = 0
with self.timer('get_batch'):
batch = next(dataloader_iterator)
if self.progress_bar is not None:
self.progress_bar.unpause()
else:
batch = None
batch_list.append(batch)
batch_step += 1
# setup accumulation
if self.train_config.gradient_accumulation_steps == -1:
@@ -1701,7 +1962,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
# flush()
### HOOK ###
loss_dict = self.hook_train_loop(batch)
with self.accelerator.accumulate(self.modules_being_trained):
loss_dict = self.hook_train_loop(batch_list)
self.timer.stop('train_loop')
if not did_first_flush:
flush()
@@ -1713,7 +1975,12 @@ class BaseSDTrainProcess(BaseTrainProcess):
with torch.no_grad():
# torch.cuda.empty_cache()
if self.train_config.optimizer.lower().startswith('dadaptation') or \
# if optimizer has get_lrs method, then use it
if hasattr(optimizer, 'get_avg_learning_rate'):
learning_rate = optimizer.get_avg_learning_rate()
elif hasattr(optimizer, 'get_learning_rates'):
learning_rate = optimizer.get_learning_rates()[0]
elif self.train_config.optimizer.lower().startswith('dadaptation') or \
self.train_config.optimizer.lower().startswith('prodigy'):
learning_rate = (
optimizer.param_groups[0]["d"] *
@@ -1726,7 +1993,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
for key, value in loss_dict.items():
prog_bar_string += f" {key}: {value:.3e}"
self.progress_bar.set_postfix_str(prog_bar_string)
if self.progress_bar is not None:
self.progress_bar.set_postfix_str(prog_bar_string)
# if the batch is a DataLoaderBatchDTO, then we need to clean it up
if isinstance(batch, DataLoaderBatchDTO):
@@ -1735,43 +2003,86 @@ class BaseSDTrainProcess(BaseTrainProcess):
# don't do on first step
if self.step_num != self.start_step:
if is_sample_step or is_save_step:
self.accelerator.wait_for_everyone()
if is_sample_step:
self.progress_bar.pause()
if self.progress_bar is not None:
self.progress_bar.pause()
flush()
# print above the progress bar
if self.train_config.free_u:
self.sd.pipeline.disable_freeu()
self.sample(self.step_num)
self.ensure_params_requires_grad()
self.progress_bar.unpause()
if self.train_config.unload_text_encoder:
# make sure the text encoder is unloaded
self.sd.text_encoder_to('cpu')
flush()
if is_save_step:
# print above the progress bar
self.progress_bar.pause()
self.print(f"Saving at step {self.step_num}")
self.save(self.step_num)
self.ensure_params_requires_grad()
self.progress_bar.unpause()
if self.logging_config.log_every and self.step_num % self.logging_config.log_every == 0:
self.progress_bar.pause()
with self.timer('log_to_tensorboard'):
# log to tensorboard
if self.writer is not None:
for key, value in loss_dict.items():
self.writer.add_scalar(f"{key}", value, self.step_num)
self.writer.add_scalar(f"lr", learning_rate, self.step_num)
if self.progress_bar is not None:
self.progress_bar.unpause()
if is_save_step:
self.accelerator
# print above the progress bar
if self.progress_bar is not None:
self.progress_bar.pause()
print_acc(f"Saving at step {self.step_num}")
self.save(self.step_num)
self.ensure_params_requires_grad()
if self.progress_bar is not None:
self.progress_bar.unpause()
if self.logging_config.log_every and self.step_num % self.logging_config.log_every == 0:
if self.progress_bar is not None:
self.progress_bar.pause()
with self.timer('log_to_tensorboard'):
# log to tensorboard
if self.accelerator.is_main_process:
if self.writer is not None:
for key, value in loss_dict.items():
self.writer.add_scalar(f"{key}", value, self.step_num)
self.writer.add_scalar(f"lr", learning_rate, self.step_num)
if self.progress_bar is not None:
self.progress_bar.unpause()
if self.accelerator.is_main_process:
# log to logger
self.logger.log({
'learning_rate': learning_rate,
})
for key, value in loss_dict.items():
self.logger.log({
f'loss/{key}': value,
})
elif self.logging_config.log_every is None:
if self.accelerator.is_main_process:
# log every step
self.logger.log({
'learning_rate': learning_rate,
})
for key, value in loss_dict.items():
self.logger.log({
f'loss/{key}': value,
})
if self.performance_log_every > 0 and self.step_num % self.performance_log_every == 0:
self.progress_bar.pause()
if self.progress_bar is not None:
self.progress_bar.pause()
# print the timers and clear them
self.timer.print()
self.timer.reset()
self.progress_bar.unpause()
if self.progress_bar is not None:
self.progress_bar.unpause()
# commit log
if self.accelerator.is_main_process:
self.logger.commit(step=self.step_num)
# sets progress bar to match out step
self.progress_bar.update(step - self.progress_bar.n)
if self.progress_bar is not None:
self.progress_bar.update(step - self.progress_bar.n)
#############################
# End of step
@@ -1785,14 +2096,20 @@ class BaseSDTrainProcess(BaseTrainProcess):
###################################################################
## END TRAIN LOOP
###################################################################
self.progress_bar.close()
self.accelerator.wait_for_everyone()
if self.progress_bar is not None:
self.progress_bar.close()
if self.train_config.free_u:
self.sd.pipeline.disable_freeu()
if not self.train_config.disable_sampling:
self.sample(self.step_num)
print("")
self.save()
self.logger.commit(step=self.step_num)
print_acc("")
if self.accelerator.is_main_process:
self.save()
self.logger.finish()
self.accelerator.end_training()
if self.save_config.push_to_hub:
if("HF_TOKEN" not in os.environ):
interpreter_login(new_session=False, write_permission=True)
@@ -1817,6 +2134,8 @@ class BaseSDTrainProcess(BaseTrainProcess):
repo_id: str,
private: bool = False,
):
if not self.accelerator.is_main_process:
return
readme_content = self._generate_readme(repo_id)
readme_path = os.path.join(self.save_root, "README.md")
with open(readme_path, "w", encoding="utf-8") as f:
@@ -1859,12 +2178,17 @@ class BaseSDTrainProcess(BaseTrainProcess):
tags.append("stable-diffusion-xl")
if self.model_config.is_flux:
tags.append("flux")
if self.model_config.is_lumina2:
tags.append("lumina2")
if self.model_config.is_v3:
tags.append("sd3")
if self.network_config:
tags.extend(
[
"lora",
"diffusers",
"template:sd-lora",
"ai-toolkit",
]
)
@@ -1899,7 +2223,7 @@ class BaseSDTrainProcess(BaseTrainProcess):
},
}
)
dtype = "torch.bfloat16" if self.model_config.is_flux else "torch.float16"
# Construct the README content
readme_content = f"""---
tags:
@@ -1921,10 +2245,25 @@ Model trained with [AI Toolkit by Ostris](https://github.com/ostris/ai-toolkit)
{"You should use `" + instance_prompt + "` to trigger the image generation." if instance_prompt else "No trigger words defined."}
## Download model
## Download model and use it with ComfyUI, AUTOMATIC1111, SD.Next, Invoke AI, etc.
Weights for this model are available in Safetensors format.
[Download](/{repo_id}/tree/main) them in the Files & versions tab.
## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers)
```py
from diffusers import AutoPipelineForText2Image
import torch
pipeline = AutoPipelineForText2Image.from_pretrained('{base_model}', torch_dtype={dtype}).to('cuda')
pipeline.load_lora_weights('{repo_id}', weight_name='{self.job.name}.safetensors')
image = pipeline('{instance_prompt if not widgets else self.sample_config.prompts[0]}').images[0]
image.save("my_image.png")
```
For more details, including weighting, merging and fusing LoRAs, check the [documentation on loading LoRAs in diffusers](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters)
"""
return readme_content

View File

@@ -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)

View File

@@ -1,7 +1,7 @@
torch
torchvision
torch==2.5.1
torchvision==0.20.1
safetensors
git+https://github.com/huggingface/diffusers.git
git+https://github.com/zhuole1025/diffusers@lumina2
transformers
lycoris-lora==1.8.3
flatten_json
@@ -13,7 +13,8 @@ invisible-watermark
einops
accelerate
toml
albumentations
albumentations==1.4.15
albucore==0.0.16
pydantic
omegaconf
k-diffusion
@@ -26,7 +27,9 @@ bitsandbytes
hf_transfer
lpips
pytorch_fid
optimum-quanto
optimum-quanto==0.2.4
sentencepiece
huggingface_hub
peft
peft
gradio
python-slugify

23
run.py
View File

@@ -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
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():
@@ -70,7 +76,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,7 +86,7 @@ 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
if not args.recover:
print_end_message(jobs_completed, jobs_failed)

View 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.")

View 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")

3
todo_multigpu.md Normal file
View File

@@ -0,0 +1,3 @@
- only do ema on main device? shouldne be needed other than saving and sampling
- check when to unwrap model and what it does
- disable timer for non main local

17
toolkit/accelerator.py Normal file
View 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

View File

@@ -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:
@@ -202,6 +206,16 @@ 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)
class EmbeddingConfig:
@@ -213,6 +227,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 +255,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 +304,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
@@ -318,8 +346,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)
@@ -358,13 +386,37 @@ 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', 0.1)
# 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)
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 +424,7 @@ 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_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 +439,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 +469,25 @@ class ModelConfig:
# only for flux for now
self.quantize = kwargs.get("quantize", False)
self.quantize_te = kwargs.get("quantize_te", self.quantize)
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", {})
if self.ignore_if_contains is not None or self.only_if_contains is not None:
if not self.is_flux:
raise ValueError("ignore_if_contains and only_if_contains are only supported with flux models currently")
# 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)
class EMAConfig:
@@ -425,6 +496,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 +589,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):
@@ -580,6 +658,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', [])
@@ -640,6 +720,7 @@ 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,
):
self.width: int = width
self.height: int = height
@@ -697,6 +778,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
@@ -757,7 +840,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
@@ -840,3 +926,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

View File

@@ -1,3 +1,4 @@
import math
import torch
import sys
@@ -13,10 +14,13 @@ 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
sys.path.append(REPOS_ROOT)
from typing import TYPE_CHECKING, Union, Iterator, Mapping, Any, Tuple, List, Optional, Dict
@@ -90,6 +94,8 @@ 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.conditional_embeds: Optional[torch.Tensor] = None
self.unconditional_embeds: Optional[torch.Tensor] = None
@@ -117,11 +123,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':
@@ -198,6 +204,9 @@ class CustomAdapter(torch.nn.Module):
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)
else:
raise ValueError(f"unknown adapter type: {self.adapter_type}")
@@ -257,6 +266,13 @@ 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 == '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 +288,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
@@ -397,12 +413,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 +429,13 @@ 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)
pass
@@ -445,8 +468,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 +479,11 @@ 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
else:
raise NotImplementedError
@@ -472,7 +500,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 == 'clip_fusion' or self.adapter_type == 'ilora' or self.adapter_type == 'vision_direct' or self.adapter_type == 'redux':
return prompt
elif self.adapter_type == 'text_encoder':
# todo allow for training
@@ -594,7 +622,7 @@ class CustomAdapter(torch.nn.Module):
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 +644,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 +725,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 +785,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 +796,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 +822,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 +931,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 +951,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}")
@@ -881,16 +978,28 @@ class CustomAdapter(torch.nn.Module):
for attn_processor in self.te_adapter.adapter_modules:
yield from attn_processor.parameters(recurse)
elif self.config.type == 'vision_direct':
for attn_processor in self.vd_adapter.adapter_modules:
yield from attn_processor.parameters(recurse)
if self.config.train_image_encoder:
yield from self.vision_encoder.parameters(recurse)
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)
else:
raise NotImplementedError
@@ -908,4 +1017,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

View File

@@ -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
@@ -90,7 +92,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 +104,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 +130,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 +142,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 +245,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(),
@@ -435,17 +437,31 @@ 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}")
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 +472,26 @@ 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())
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")
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.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 +501,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,7 +510,7 @@ 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")
print_acc(f" - Found {len(self.file_list)} images after adding flips")
self.setup_epoch()

View File

@@ -44,19 +44,25 @@ class FileItemDTO(
self.path = kwargs.get('path', '')
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
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]
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)

View File

@@ -1,4 +1,5 @@
import base64
import glob
import hashlib
import json
import math
@@ -12,16 +13,19 @@ 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 +34,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 +125,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 +138,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 +267,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 +279,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:
@@ -438,8 +451,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 +466,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 +486,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 +505,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:
@@ -558,8 +571,8 @@ class ControlFileItemDTOMixin:
img = Image.open(self.control_path).convert('RGB')
img = exif_transpose(img)
except Exception as e:
print(f"Error: {e}")
print(f"Error loading image: {self.control_path}")
print_acc(f"Error: {e}")
print_acc(f"Error loading image: {self.control_path}")
if self.full_size_control_images:
# we just scale them to 512x512:
@@ -629,11 +642,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 +659,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 +749,27 @@ 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.is_vision_clip_cached:
self.clip_image_embeds = load_file(self.get_clip_vision_embeddings_path())
@@ -742,14 +779,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 +797,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 +985,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 +1003,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 +1089,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 +1170,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 +1246,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 +1256,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 +1283,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 +1377,89 @@ 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
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 = '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
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 +1474,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 +1517,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
View 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

View File

@@ -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,

View File

@@ -481,9 +481,9 @@ 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
# 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])

View File

@@ -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
View 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()

View File

@@ -63,7 +63,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 +163,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,
@@ -223,6 +224,7 @@ 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":
@@ -232,7 +234,7 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
self.peft_format = peft_format
# always do peft for flux only for now
if self.is_flux:
if self.is_flux or self.is_v3 or self.is_lumina2:
self.peft_format = True
if self.peft_format:
@@ -273,7 +275,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:
unet_prefix = f"lora_transformer"
if self.peft_format:
unet_prefix = "transformer"
@@ -305,15 +307,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,10 +328,16 @@ 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
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]):
if self.only_if_contains is not None and not any([word in clean_name for word in self.only_if_contains]):
continue
dim = None
@@ -428,6 +436,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)

View File

@@ -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(),

View 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

View File

@@ -0,0 +1,356 @@
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.detach() for x in get_lpips_features(
tensors_n1p1.to(device, dtype=torch.float32))]
return lpips_feat_list
def forward(
self,
noise_pred,
noisy_latents,
timesteps,
batch: DataLoaderBatchDTO,
scheduler: CustomFlowMatchEulerDiscreteScheduler,
lpips_weight=20.0,
clip_weight=0.1,
pixel_weight=1.0
):
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)
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
pred_clip_output = self.get_siglip_features(pred_images)
lpips_feat_list_pred = self.get_lpips_features(pred_images.float())
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
target_clip_output = self.get_siglip_features(target_img).detach()
lpips_feat_list_target = self.get_lpips_features(target_img.float())
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
lpips_loss = 0
for idx, lpips_feat in enumerate(lpips_feat_list_pred):
lpips_loss += torch.nn.functional.mse_loss(
lpips_feat.float(), lpips_feat_list_target[idx].float()
) * lpips_weight
if 'lpips_loss' not in self.losses:
self.losses['lpips_loss'] = lpips_loss.item()
else:
self.losses['lpips_loss'] += 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

176
toolkit/models/flux.py Normal file
View 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)

View 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

View File

@@ -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)

560
toolkit/models/lumina2.py Normal file
View File

@@ -0,0 +1,560 @@
# 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]:
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)

View 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
View 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

View File

@@ -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

View File

@@ -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

View File

@@ -28,6 +28,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 +53,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 +85,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 +95,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

View 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

View 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'])

View File

@@ -0,0 +1,335 @@
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
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)

View 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

View 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

View File

@@ -14,6 +14,7 @@ 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
if is_torch_xla_available():
@@ -1235,6 +1236,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 +1285,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 +1362,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,6 +1418,7 @@ class FluxWithCFGPipeline(FluxPipeline):
# Offload all models
self.maybe_free_model_hooks()
restore_flux_guidance(self.transformer)
if not return_dict:
return (image,)

6
toolkit/print.py Normal file
View File

@@ -0,0 +1,6 @@
from toolkit.accelerator import get_accelerator
def print_acc(*args, **kwargs):
if get_accelerator().is_local_main_process:
print(*args, **kwargs)

View File

@@ -76,6 +76,35 @@ 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
}
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 +149,13 @@ 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
}
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")

View File

@@ -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,29 @@ 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]
# 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
@@ -79,12 +104,13 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
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):
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,63 @@ class CustomFlowMatchEulerDiscreteScheduler(FlowMatchEulerDiscreteScheduler):
self.timesteps = timesteps.to(device=device)
return timesteps
elif timestep_type == 'flux_shift' or timestep_type == 'lumina2_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 latents is None:
raise ValueError('latents is None')
h = latents.shape[2] // 2 # Divide by ph
w = latents.shape[3] // 2 # Divide by pw
image_seq_len = h * w
# todo need to know the mu for the shift
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)
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
timesteps = sigmas * self.config.num_train_timesteps
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}")

View File

@@ -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

View File

@@ -517,6 +517,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 +569,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

View File

@@ -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