Compare commits
93 Commits
60232def91
...
lumina2
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ed1deb71c4 | ||
|
|
4de6a825fa | ||
|
|
9a7266275d | ||
|
|
d138f07365 | ||
|
|
c6d8eedb94 | ||
|
|
af5e760be1 | ||
|
|
ff3d54bb5b | ||
|
|
0e75724b4d | ||
|
|
376bb1bf6f | ||
|
|
216ab164ce | ||
|
|
e6180d1e1d | ||
|
|
15a57bc89f | ||
|
|
e5355bf8d5 | ||
|
|
34a1c6947a | ||
|
|
2141c6e06c | ||
|
|
1188cf1e8a | ||
|
|
5e663746b8 | ||
|
|
441474e81f | ||
|
|
a6a690f796 | ||
|
|
6191f19e55 | ||
|
|
bbfba0c188 | ||
|
|
e1549ad54d | ||
|
|
04abe57c76 | ||
|
|
89dd041b97 | ||
|
|
29122b1a54 | ||
|
|
6a8e3d8610 | ||
|
|
4c8a9e1b88 | ||
|
|
fadb2f3a76 | ||
|
|
4723f23c0d | ||
|
|
8ef07a9c36 | ||
|
|
92ce93140e | ||
|
|
f213996aa5 | ||
|
|
cbe31eaf0a | ||
|
|
67c2e44edb | ||
|
|
96d418bb95 | ||
|
|
894374b2e9 | ||
|
|
6509ba4484 | ||
|
|
025ee3dd3d | ||
|
|
58f9d01c2b | ||
|
|
e72b59a8e9 | ||
|
|
4aa19b5c1d | ||
|
|
4747716867 | ||
|
|
22cd40d7b9 | ||
|
|
3400882a80 | ||
|
|
9f94c7b61e | ||
|
|
bedb8197a2 | ||
|
|
e3ebd73610 | ||
|
|
dd931757cd | ||
|
|
0640cdf569 | ||
|
|
0b048d0dde | ||
|
|
473d455f44 | ||
|
|
ce759ebd8c | ||
|
|
628a7923a3 | ||
|
|
3922981996 | ||
|
|
ab22674980 | ||
|
|
9452929300 | ||
|
|
a800c9d19e | ||
|
|
28e6f00790 | ||
|
|
67e0aca750 | ||
|
|
f05224970f | ||
|
|
b4f64de4c2 | ||
|
|
2e5f6668dc | ||
|
|
e4c82803e1 | ||
|
|
69aa92bce5 | ||
|
|
a508caad1d | ||
|
|
58537fc92b | ||
|
|
86b5938cf3 | ||
|
|
6b4034122f | ||
|
|
10817696fb | ||
|
|
037ce11740 | ||
|
|
04424fe2d6 | ||
|
|
40a8ff5731 | ||
|
|
2776221497 | ||
|
|
f85ad452c6 | ||
|
|
dd889086f4 | ||
|
|
bc693488eb | ||
|
|
d97c55cd96 | ||
|
|
79b4e04b80 | ||
|
|
951e223481 | ||
|
|
fc34a69bec | ||
|
|
279ee65177 | ||
|
|
3a1f464132 | ||
|
|
5c8fcc8a4e | ||
|
|
121a760c19 | ||
|
|
e5fadddd45 | ||
|
|
d44d4eb61a | ||
|
|
7d9ab22405 | ||
|
|
7ed8c51f20 | ||
|
|
6df33156f0 | ||
|
|
40f5c59da0 | ||
|
|
3e71a99df0 | ||
|
|
562405923f | ||
|
|
f84bd6d7a6 |
6
.gitignore
vendored
6
.gitignore
vendored
@@ -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
28
.vscode/launch.json
vendored
Normal file
@@ -0,0 +1,28 @@
|
||||
{
|
||||
"version": "0.2.0",
|
||||
"configurations": [
|
||||
{
|
||||
"name": "Run current config",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${workspaceFolder}/run.py",
|
||||
"args": [
|
||||
"${file}"
|
||||
],
|
||||
"env": {
|
||||
"CUDA_LAUNCH_BLOCKING": "1",
|
||||
"DEBUG_TOOLKIT": "1"
|
||||
},
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
{
|
||||
"name": "Python: Debug Current File",
|
||||
"type": "python",
|
||||
"request": "launch",
|
||||
"program": "${file}",
|
||||
"console": "integratedTerminal",
|
||||
"justMyCode": false
|
||||
},
|
||||
]
|
||||
}
|
||||
62
README.md
62
README.md
@@ -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
|
||||

|
||||
|
||||
|
||||
## 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
BIN
assets/lora_ease_ui.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 340 KiB |
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
107
config/examples/train_full_fine_tune_flex.yaml
Normal file
@@ -0,0 +1,107 @@
|
||||
---
|
||||
# This configuration requires 48GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_finetune_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lognorm_blend'
|
||||
timestep_type: 'sigmoid'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true # flex is flux architecture
|
||||
# full finetuning quantized models is a crapshoot and results in subpar outputs
|
||||
# quantize: true
|
||||
# you can quantize just the T5 text encoder here to save vram
|
||||
quantize_te: true
|
||||
# only train the transformer blocks
|
||||
only_if_contains:
|
||||
- "transformer.transformer_blocks."
|
||||
- "transformer.single_transformer_blocks."
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
99
config/examples/train_full_fine_tune_lumina.yaml
Normal file
@@ -0,0 +1,99 @@
|
||||
---
|
||||
# This configuration requires 24GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_finetune_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adafactor"
|
||||
lr: 3e-5
|
||||
|
||||
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
|
||||
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
|
||||
|
||||
# do_paramiter_swapping: true
|
||||
# paramiter_swapping_factor: 0.9
|
||||
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
# ema_config:
|
||||
# use_ema: true
|
||||
# ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
101
config/examples/train_lora_flex_24gb.yaml
Normal file
101
config/examples/train_lora_flex_24gb.yaml
Normal file
@@ -0,0 +1,101 @@
|
||||
---
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_flex_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
|
||||
bypass_guidance_embedding: true
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with flex
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "ostris/Flex.1-alpha"
|
||||
is_flux: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
quantize_kwargs:
|
||||
exclude:
|
||||
- "*time_text_embed*" # exclude the time text embedder from quantization
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: "" # not used on flex
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
96
config/examples/train_lora_lumina.yaml
Normal file
96
config/examples/train_lora_lumina.yaml
Normal file
@@ -0,0 +1,96 @@
|
||||
---
|
||||
# This configuration requires 20GB of VRAM or more to operate
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_lumina_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: bf16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 2 # how many intermittent saves to keep
|
||||
save_format: 'diffusers' # 'diffusers'
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
# cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
|
||||
train:
|
||||
batch_size: 1
|
||||
|
||||
# can be 'sigmoid', 'linear', or 'lumina2_shift'
|
||||
timestep_type: 'lumina2_shift'
|
||||
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # probably won't work with lumina2
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch" # for training only
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
|
||||
is_lumina2: true # lumina2 architecture
|
||||
# you can quantize just the Gemma2 text encoder here to save vram
|
||||
quantize_te: true
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4.0
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
97
config/examples/train_lora_sd35_large_24gb.yaml
Normal file
@@ -0,0 +1,97 @@
|
||||
---
|
||||
# NOTE!! THIS IS CURRENTLY EXPERIMENTAL AND UNDER DEVELOPMENT. SOME THINGS WILL CHANGE
|
||||
job: extension
|
||||
config:
|
||||
# this name will be the folder and filename name
|
||||
name: "my_first_sd3l_lora_v1"
|
||||
process:
|
||||
- type: 'sd_trainer'
|
||||
# root folder to save training sessions/samples/weights
|
||||
training_folder: "output"
|
||||
# uncomment to see performance stats in the terminal every N steps
|
||||
# performance_log_every: 1000
|
||||
device: cuda:0
|
||||
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
|
||||
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
|
||||
# trigger_word: "p3r5on"
|
||||
network:
|
||||
type: "lora"
|
||||
linear: 16
|
||||
linear_alpha: 16
|
||||
save:
|
||||
dtype: float16 # precision to save
|
||||
save_every: 250 # save every this many steps
|
||||
max_step_saves_to_keep: 4 # how many intermittent saves to keep
|
||||
push_to_hub: false #change this to True to push your trained model to Hugging Face.
|
||||
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
|
||||
# hf_repo_id: your-username/your-model-slug
|
||||
# hf_private: true #whether the repo is private or public
|
||||
datasets:
|
||||
# datasets are a folder of images. captions need to be txt files with the same name as the image
|
||||
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
|
||||
# images will automatically be resized and bucketed into the resolution specified
|
||||
# on windows, escape back slashes with another backslash so
|
||||
# "C:\\path\\to\\images\\folder"
|
||||
- folder_path: "/path/to/images/folder"
|
||||
caption_ext: "txt"
|
||||
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
|
||||
shuffle_tokens: false # shuffle caption order, split by commas
|
||||
cache_latents_to_disk: true # leave this true unless you know what you're doing
|
||||
resolution: [ 1024 ]
|
||||
train:
|
||||
batch_size: 1
|
||||
steps: 2000 # total number of steps to train 500 - 4000 is a good range
|
||||
gradient_accumulation_steps: 1
|
||||
train_unet: true
|
||||
train_text_encoder: false # May not fully work with SD3 yet
|
||||
gradient_checkpointing: true # need the on unless you have a ton of vram
|
||||
noise_scheduler: "flowmatch"
|
||||
timestep_type: "linear" # linear or sigmoid
|
||||
optimizer: "adamw8bit"
|
||||
lr: 1e-4
|
||||
# uncomment this to skip the pre training sample
|
||||
# skip_first_sample: true
|
||||
# uncomment to completely disable sampling
|
||||
# disable_sampling: true
|
||||
# uncomment to use new vell curved weighting. Experimental but may produce better results
|
||||
# linear_timesteps: true
|
||||
|
||||
# ema will smooth out learning, but could slow it down. Recommended to leave on.
|
||||
ema_config:
|
||||
use_ema: true
|
||||
ema_decay: 0.99
|
||||
|
||||
# will probably need this if gpu supports it for sd3, other dtypes may not work correctly
|
||||
dtype: bf16
|
||||
model:
|
||||
# huggingface model name or path
|
||||
name_or_path: "stabilityai/stable-diffusion-3.5-large"
|
||||
is_v3: true
|
||||
quantize: true # run 8bit mixed precision
|
||||
sample:
|
||||
sampler: "flowmatch" # must match train.noise_scheduler
|
||||
sample_every: 250 # sample every this many steps
|
||||
width: 1024
|
||||
height: 1024
|
||||
prompts:
|
||||
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
|
||||
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
|
||||
- "woman with red hair, playing chess at the park, bomb going off in the background"
|
||||
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
|
||||
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
|
||||
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
|
||||
- "a bear building a log cabin in the snow covered mountains"
|
||||
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
|
||||
- "hipster man with a beard, building a chair, in a wood shop"
|
||||
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
|
||||
- "a man holding a sign that says, 'this is a sign'"
|
||||
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
|
||||
neg: ""
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
guidance_scale: 4
|
||||
sample_steps: 25
|
||||
# you can add any additional meta info here. [name] is replaced with config name at top
|
||||
meta:
|
||||
name: "[name]"
|
||||
version: '1.0'
|
||||
@@ -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"]
|
||||
@@ -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
414
flux_train_ui.py
Normal file
@@ -0,0 +1,414 @@
|
||||
import os
|
||||
from huggingface_hub import whoami
|
||||
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
|
||||
import sys
|
||||
|
||||
# Add the current working directory to the Python path
|
||||
sys.path.insert(0, os.getcwd())
|
||||
|
||||
import gradio as gr
|
||||
from PIL import Image
|
||||
import torch
|
||||
import uuid
|
||||
import os
|
||||
import shutil
|
||||
import json
|
||||
import yaml
|
||||
from slugify import slugify
|
||||
from transformers import AutoProcessor, AutoModelForCausalLM
|
||||
|
||||
sys.path.insert(0, "ai-toolkit")
|
||||
from toolkit.job import get_job
|
||||
|
||||
MAX_IMAGES = 150
|
||||
|
||||
def load_captioning(uploaded_files, concept_sentence):
|
||||
uploaded_images = [file for file in uploaded_files if not file.endswith('.txt')]
|
||||
txt_files = [file for file in uploaded_files if file.endswith('.txt')]
|
||||
txt_files_dict = {os.path.splitext(os.path.basename(txt_file))[0]: txt_file for txt_file in txt_files}
|
||||
updates = []
|
||||
if len(uploaded_images) <= 1:
|
||||
raise gr.Error(
|
||||
"Please upload at least 2 images to train your model (the ideal number with default settings is between 4-30)"
|
||||
)
|
||||
elif len(uploaded_images) > MAX_IMAGES:
|
||||
raise gr.Error(f"For now, only {MAX_IMAGES} or less images are allowed for training")
|
||||
# Update for the captioning_area
|
||||
# for _ in range(3):
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update visibility and image for each captioning row and image
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
# Determine if the current row and image should be visible
|
||||
visible = i <= len(uploaded_images)
|
||||
|
||||
# Update visibility of the captioning row
|
||||
updates.append(gr.update(visible=visible))
|
||||
|
||||
# Update for image component - display image if available, otherwise hide
|
||||
image_value = uploaded_images[i - 1] if visible else None
|
||||
updates.append(gr.update(value=image_value, visible=visible))
|
||||
|
||||
corresponding_caption = False
|
||||
if(image_value):
|
||||
base_name = os.path.splitext(os.path.basename(image_value))[0]
|
||||
print(base_name)
|
||||
print(image_value)
|
||||
if base_name in txt_files_dict:
|
||||
print("entrou")
|
||||
with open(txt_files_dict[base_name], 'r') as file:
|
||||
corresponding_caption = file.read()
|
||||
|
||||
# Update value of captioning area
|
||||
text_value = corresponding_caption if visible and corresponding_caption else "[trigger]" if visible and concept_sentence else None
|
||||
updates.append(gr.update(value=text_value, visible=visible))
|
||||
|
||||
# Update for the sample caption area
|
||||
updates.append(gr.update(visible=True))
|
||||
# Update prompt samples
|
||||
updates.append(gr.update(placeholder=f'A portrait of person in a bustling cafe {concept_sentence}', value=f'A person in a bustling cafe {concept_sentence}'))
|
||||
updates.append(gr.update(placeholder=f"A mountainous landscape in the style of {concept_sentence}"))
|
||||
updates.append(gr.update(placeholder=f"A {concept_sentence} in a mall"))
|
||||
updates.append(gr.update(visible=True))
|
||||
return updates
|
||||
|
||||
def hide_captioning():
|
||||
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
|
||||
|
||||
def create_dataset(*inputs):
|
||||
print("Creating dataset")
|
||||
images = inputs[0]
|
||||
destination_folder = str(f"datasets/{uuid.uuid4()}")
|
||||
if not os.path.exists(destination_folder):
|
||||
os.makedirs(destination_folder)
|
||||
|
||||
jsonl_file_path = os.path.join(destination_folder, "metadata.jsonl")
|
||||
with open(jsonl_file_path, "a") as jsonl_file:
|
||||
for index, image in enumerate(images):
|
||||
new_image_path = shutil.copy(image, destination_folder)
|
||||
|
||||
original_caption = inputs[index + 1]
|
||||
file_name = os.path.basename(new_image_path)
|
||||
|
||||
data = {"file_name": file_name, "prompt": original_caption}
|
||||
|
||||
jsonl_file.write(json.dumps(data) + "\n")
|
||||
|
||||
return destination_folder
|
||||
|
||||
|
||||
def run_captioning(images, concept_sentence, *captions):
|
||||
#Load internally to not consume resources for training
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
torch_dtype = torch.float16
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
"multimodalart/Florence-2-large-no-flash-attn", torch_dtype=torch_dtype, trust_remote_code=True
|
||||
).to(device)
|
||||
processor = AutoProcessor.from_pretrained("multimodalart/Florence-2-large-no-flash-attn", trust_remote_code=True)
|
||||
|
||||
captions = list(captions)
|
||||
for i, image_path in enumerate(images):
|
||||
print(captions[i])
|
||||
if isinstance(image_path, str): # If image is a file path
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
|
||||
prompt = "<DETAILED_CAPTION>"
|
||||
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
|
||||
|
||||
generated_ids = model.generate(
|
||||
input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=1024, num_beams=3
|
||||
)
|
||||
|
||||
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
|
||||
parsed_answer = processor.post_process_generation(
|
||||
generated_text, task=prompt, image_size=(image.width, image.height)
|
||||
)
|
||||
caption_text = parsed_answer["<DETAILED_CAPTION>"].replace("The image shows ", "")
|
||||
if concept_sentence:
|
||||
caption_text = f"{caption_text} [trigger]"
|
||||
captions[i] = caption_text
|
||||
|
||||
yield captions
|
||||
model.to("cpu")
|
||||
del model
|
||||
del processor
|
||||
|
||||
def recursive_update(d, u):
|
||||
for k, v in u.items():
|
||||
if isinstance(v, dict) and v:
|
||||
d[k] = recursive_update(d.get(k, {}), v)
|
||||
else:
|
||||
d[k] = v
|
||||
return d
|
||||
|
||||
def start_training(
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options,
|
||||
):
|
||||
push_to_hub = True
|
||||
if not lora_name:
|
||||
raise gr.Error("You forgot to insert your LoRA name! This name has to be unique.")
|
||||
try:
|
||||
if whoami()["auth"]["accessToken"]["role"] == "write" or "repo.write" in whoami()["auth"]["accessToken"]["fineGrained"]["scoped"][0]["permissions"]:
|
||||
gr.Info(f"Starting training locally {whoami()['name']}. Your LoRA will be available locally and in Hugging Face after it finishes.")
|
||||
else:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
except:
|
||||
push_to_hub = False
|
||||
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
|
||||
|
||||
print("Started training")
|
||||
slugged_lora_name = slugify(lora_name)
|
||||
|
||||
# Load the default config
|
||||
with open("config/examples/train_lora_flux_24gb.yaml", "r") as f:
|
||||
config = yaml.safe_load(f)
|
||||
|
||||
# Update the config with user inputs
|
||||
config["config"]["name"] = slugged_lora_name
|
||||
config["config"]["process"][0]["model"]["low_vram"] = low_vram
|
||||
config["config"]["process"][0]["train"]["skip_first_sample"] = True
|
||||
config["config"]["process"][0]["train"]["steps"] = int(steps)
|
||||
config["config"]["process"][0]["train"]["lr"] = float(lr)
|
||||
config["config"]["process"][0]["network"]["linear"] = int(rank)
|
||||
config["config"]["process"][0]["network"]["linear_alpha"] = int(rank)
|
||||
config["config"]["process"][0]["datasets"][0]["folder_path"] = dataset_folder
|
||||
config["config"]["process"][0]["save"]["push_to_hub"] = push_to_hub
|
||||
if(push_to_hub):
|
||||
try:
|
||||
username = whoami()["name"]
|
||||
except:
|
||||
raise gr.Error("Error trying to retrieve your username. Are you sure you are logged in with Hugging Face?")
|
||||
config["config"]["process"][0]["save"]["hf_repo_id"] = f"{username}/{slugged_lora_name}"
|
||||
config["config"]["process"][0]["save"]["hf_private"] = True
|
||||
if concept_sentence:
|
||||
config["config"]["process"][0]["trigger_word"] = concept_sentence
|
||||
|
||||
if sample_1 or sample_2 or sample_3:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = False
|
||||
config["config"]["process"][0]["sample"]["sample_every"] = steps
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 28
|
||||
config["config"]["process"][0]["sample"]["prompts"] = []
|
||||
if sample_1:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_1)
|
||||
if sample_2:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_2)
|
||||
if sample_3:
|
||||
config["config"]["process"][0]["sample"]["prompts"].append(sample_3)
|
||||
else:
|
||||
config["config"]["process"][0]["train"]["disable_sampling"] = True
|
||||
if(model_to_train == "schnell"):
|
||||
config["config"]["process"][0]["model"]["name_or_path"] = "black-forest-labs/FLUX.1-schnell"
|
||||
config["config"]["process"][0]["model"]["assistant_lora_path"] = "ostris/FLUX.1-schnell-training-adapter"
|
||||
config["config"]["process"][0]["sample"]["sample_steps"] = 4
|
||||
if(use_more_advanced_options):
|
||||
more_advanced_options_dict = yaml.safe_load(more_advanced_options)
|
||||
config["config"]["process"][0] = recursive_update(config["config"]["process"][0], more_advanced_options_dict)
|
||||
print(config)
|
||||
|
||||
# Save the updated config
|
||||
# generate a random name for the config
|
||||
random_config_name = str(uuid.uuid4())
|
||||
os.makedirs("tmp", exist_ok=True)
|
||||
config_path = f"tmp/{random_config_name}-{slugged_lora_name}.yaml"
|
||||
with open(config_path, "w") as f:
|
||||
yaml.dump(config, f)
|
||||
|
||||
# run the job locally
|
||||
job = get_job(config_path)
|
||||
job.run()
|
||||
job.cleanup()
|
||||
|
||||
return f"Training completed successfully. Model saved as {slugged_lora_name}"
|
||||
|
||||
config_yaml = '''
|
||||
device: cuda:0
|
||||
model:
|
||||
is_flux: true
|
||||
quantize: true
|
||||
network:
|
||||
linear: 16 #it will overcome the 'rank' parameter
|
||||
linear_alpha: 16 #you can have an alpha different than the ranking if you'd like
|
||||
type: lora
|
||||
sample:
|
||||
guidance_scale: 3.5
|
||||
height: 1024
|
||||
neg: '' #doesn't work for FLUX
|
||||
sample_every: 1000
|
||||
sample_steps: 28
|
||||
sampler: flowmatch
|
||||
seed: 42
|
||||
walk_seed: true
|
||||
width: 1024
|
||||
save:
|
||||
dtype: float16
|
||||
hf_private: true
|
||||
max_step_saves_to_keep: 4
|
||||
push_to_hub: true
|
||||
save_every: 10000
|
||||
train:
|
||||
batch_size: 1
|
||||
dtype: bf16
|
||||
ema_config:
|
||||
ema_decay: 0.99
|
||||
use_ema: true
|
||||
gradient_accumulation_steps: 1
|
||||
gradient_checkpointing: true
|
||||
noise_scheduler: flowmatch
|
||||
optimizer: adamw8bit #options: prodigy, dadaptation, adamw, adamw8bit, lion, lion8bit
|
||||
train_text_encoder: false #probably doesn't work for flux
|
||||
train_unet: true
|
||||
'''
|
||||
|
||||
theme = gr.themes.Monochrome(
|
||||
text_size=gr.themes.Size(lg="18px", md="15px", sm="13px", xl="22px", xs="12px", xxl="24px", xxs="9px"),
|
||||
font=[gr.themes.GoogleFont("Source Sans Pro"), "ui-sans-serif", "system-ui", "sans-serif"],
|
||||
)
|
||||
css = """
|
||||
h1{font-size: 2em}
|
||||
h3{margin-top: 0}
|
||||
#component-1{text-align:center}
|
||||
.main_ui_logged_out{opacity: 0.3; pointer-events: none}
|
||||
.tabitem{border: 0px}
|
||||
.group_padding{padding: .55em}
|
||||
"""
|
||||
with gr.Blocks(theme=theme, css=css) as demo:
|
||||
gr.Markdown(
|
||||
"""# LoRA Ease for FLUX 🧞♂️
|
||||
### Train a high quality FLUX LoRA in a breeze ༄ using [Ostris' AI Toolkit](https://github.com/ostris/ai-toolkit)"""
|
||||
)
|
||||
with gr.Column() as main_ui:
|
||||
with gr.Row():
|
||||
lora_name = gr.Textbox(
|
||||
label="The name of your LoRA",
|
||||
info="This has to be a unique name",
|
||||
placeholder="e.g.: Persian Miniature Painting style, Cat Toy",
|
||||
)
|
||||
concept_sentence = gr.Textbox(
|
||||
label="Trigger word/sentence",
|
||||
info="Trigger word or sentence to be used",
|
||||
placeholder="uncommon word like p3rs0n or trtcrd, or sentence like 'in the style of CNSTLL'",
|
||||
interactive=True,
|
||||
)
|
||||
with gr.Group(visible=True) as image_upload:
|
||||
with gr.Row():
|
||||
images = gr.File(
|
||||
file_types=["image", ".txt"],
|
||||
label="Upload your images",
|
||||
file_count="multiple",
|
||||
interactive=True,
|
||||
visible=True,
|
||||
scale=1,
|
||||
)
|
||||
with gr.Column(scale=3, visible=False) as captioning_area:
|
||||
with gr.Column():
|
||||
gr.Markdown(
|
||||
"""# Custom captioning
|
||||
<p style="margin-top:0">You can optionally add a custom caption for each image (or use an AI model for this). [trigger] will represent your concept sentence/trigger word.</p>
|
||||
""", elem_classes="group_padding")
|
||||
do_captioning = gr.Button("Add AI captions with Florence-2")
|
||||
output_components = [captioning_area]
|
||||
caption_list = []
|
||||
for i in range(1, MAX_IMAGES + 1):
|
||||
locals()[f"captioning_row_{i}"] = gr.Row(visible=False)
|
||||
with locals()[f"captioning_row_{i}"]:
|
||||
locals()[f"image_{i}"] = gr.Image(
|
||||
type="filepath",
|
||||
width=111,
|
||||
height=111,
|
||||
min_width=111,
|
||||
interactive=False,
|
||||
scale=2,
|
||||
show_label=False,
|
||||
show_share_button=False,
|
||||
show_download_button=False,
|
||||
)
|
||||
locals()[f"caption_{i}"] = gr.Textbox(
|
||||
label=f"Caption {i}", scale=15, interactive=True
|
||||
)
|
||||
|
||||
output_components.append(locals()[f"captioning_row_{i}"])
|
||||
output_components.append(locals()[f"image_{i}"])
|
||||
output_components.append(locals()[f"caption_{i}"])
|
||||
caption_list.append(locals()[f"caption_{i}"])
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
steps = gr.Number(label="Steps", value=1000, minimum=1, maximum=10000, step=1)
|
||||
lr = gr.Number(label="Learning Rate", value=4e-4, minimum=1e-6, maximum=1e-3, step=1e-6)
|
||||
rank = gr.Number(label="LoRA Rank", value=16, minimum=4, maximum=128, step=4)
|
||||
model_to_train = gr.Radio(["dev", "schnell"], value="dev", label="Model to train")
|
||||
low_vram = gr.Checkbox(label="Low VRAM", value=True)
|
||||
with gr.Accordion("Even more advanced options", open=False):
|
||||
use_more_advanced_options = gr.Checkbox(label="Use more advanced options", value=False)
|
||||
more_advanced_options = gr.Code(config_yaml, language="yaml")
|
||||
|
||||
with gr.Accordion("Sample prompts (optional)", visible=False) as sample:
|
||||
gr.Markdown(
|
||||
"Include sample prompts to test out your trained model. Don't forget to include your trigger word/sentence (optional)"
|
||||
)
|
||||
sample_1 = gr.Textbox(label="Test prompt 1")
|
||||
sample_2 = gr.Textbox(label="Test prompt 2")
|
||||
sample_3 = gr.Textbox(label="Test prompt 3")
|
||||
|
||||
output_components.append(sample)
|
||||
output_components.append(sample_1)
|
||||
output_components.append(sample_2)
|
||||
output_components.append(sample_3)
|
||||
start = gr.Button("Start training", visible=False)
|
||||
output_components.append(start)
|
||||
progress_area = gr.Markdown("")
|
||||
|
||||
dataset_folder = gr.State()
|
||||
|
||||
images.upload(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.delete(
|
||||
load_captioning,
|
||||
inputs=[images, concept_sentence],
|
||||
outputs=output_components
|
||||
)
|
||||
|
||||
images.clear(
|
||||
hide_captioning,
|
||||
outputs=[captioning_area, sample, start]
|
||||
)
|
||||
|
||||
start.click(fn=create_dataset, inputs=[images] + caption_list, outputs=dataset_folder).then(
|
||||
fn=start_training,
|
||||
inputs=[
|
||||
lora_name,
|
||||
concept_sentence,
|
||||
steps,
|
||||
lr,
|
||||
rank,
|
||||
model_to_train,
|
||||
low_vram,
|
||||
dataset_folder,
|
||||
sample_1,
|
||||
sample_2,
|
||||
sample_3,
|
||||
use_more_advanced_options,
|
||||
more_advanced_options
|
||||
],
|
||||
outputs=progress_area,
|
||||
)
|
||||
|
||||
do_captioning.click(fn=run_captioning, inputs=[images, concept_sentence] + caption_list, outputs=caption_list)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.launch(share=True, show_error=True)
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
23
run.py
@@ -20,20 +20,26 @@ if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
|
||||
torch.autograd.set_detect_anomaly(True)
|
||||
import argparse
|
||||
from toolkit.job import get_job
|
||||
from toolkit.accelerator import get_accelerator
|
||||
from toolkit.print import print_acc
|
||||
|
||||
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)
|
||||
|
||||
426
scripts/convert_diffusers_to_comfy.py
Normal file
426
scripts/convert_diffusers_to_comfy.py
Normal file
@@ -0,0 +1,426 @@
|
||||
#######################################################
|
||||
# Convert Diffusers Flux/Flex to all in one ComfyUI safetensors file
|
||||
# The VAE, T5 and clip will all be in the safetensors file
|
||||
# T5 will always be 8bit with the all in one file
|
||||
# You can save the transformer weights as bf16 or 8-bit with the --do_8_bit flag
|
||||
#
|
||||
# Download a reference model from Huggingface
|
||||
# https://huggingface.co/Comfy-Org/flux1-dev/blob/main/flux1-dev-fp8.safetensors
|
||||
#
|
||||
# Call like this for 8-bit transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors --do_8_bit
|
||||
#
|
||||
# Call like this for bf16 transformer weights:
|
||||
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors
|
||||
#
|
||||
#######################################################
|
||||
|
||||
|
||||
import argparse
|
||||
from datetime import date
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import safetensors
|
||||
import safetensors.torch
|
||||
import torch
|
||||
import tqdm
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("diffusers_path", type=str,
|
||||
help="Path to the original Flux diffusers folder.")
|
||||
parser.add_argument("quantized_state_dict_path", type=str,
|
||||
help="Path to the ComfyUI all in one template file.")
|
||||
parser.add_argument("flux_path", type=str,
|
||||
help="Output path for the Flux safetensors file.")
|
||||
parser.add_argument("--do_8_bit", action="store_true",
|
||||
help="Use 8-bit weights instead of bf16.")
|
||||
args = parser.parse_args()
|
||||
|
||||
flux_path = Path(args.flux_path)
|
||||
diffusers_path = Path(args.diffusers_path, "transformer")
|
||||
quantized_state_dict_path = Path(args.quantized_state_dict_path)
|
||||
|
||||
do_8_bit = args.do_8_bit
|
||||
|
||||
if not os.path.exists(flux_path.parent):
|
||||
os.makedirs(flux_path.parent)
|
||||
|
||||
if not diffusers_path.exists():
|
||||
print(f"Error: Missing transformer folder: {diffusers_path}")
|
||||
exit()
|
||||
|
||||
original_json_path = Path.joinpath(
|
||||
diffusers_path, "diffusion_pytorch_model.safetensors.index.json")
|
||||
if not original_json_path.exists():
|
||||
print(f"Error: Missing transformer index json: {original_json_path}")
|
||||
exit()
|
||||
|
||||
if not os.path.exists(quantized_state_dict_path):
|
||||
print(
|
||||
f"Error: Missing quantized state dict file: {args.quantized_state_dict_path}")
|
||||
exit()
|
||||
|
||||
with open(original_json_path, "r", encoding="utf-8") as f:
|
||||
original_json = json.load(f)
|
||||
|
||||
diffusers_map = {
|
||||
"time_in.in_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_1.weight",
|
||||
],
|
||||
"time_in.in_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_1.bias",
|
||||
],
|
||||
"time_in.out_layer.weight": [
|
||||
"time_text_embed.timestep_embedder.linear_2.weight",
|
||||
],
|
||||
"time_in.out_layer.bias": [
|
||||
"time_text_embed.timestep_embedder.linear_2.bias",
|
||||
],
|
||||
"vector_in.in_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_1.weight",
|
||||
],
|
||||
"vector_in.in_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_1.bias",
|
||||
],
|
||||
"vector_in.out_layer.weight": [
|
||||
"time_text_embed.text_embedder.linear_2.weight",
|
||||
],
|
||||
"vector_in.out_layer.bias": [
|
||||
"time_text_embed.text_embedder.linear_2.bias",
|
||||
],
|
||||
"guidance_in.in_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_1.weight",
|
||||
],
|
||||
"guidance_in.in_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_1.bias",
|
||||
],
|
||||
"guidance_in.out_layer.weight": [
|
||||
"time_text_embed.guidance_embedder.linear_2.weight",
|
||||
],
|
||||
"guidance_in.out_layer.bias": [
|
||||
"time_text_embed.guidance_embedder.linear_2.bias",
|
||||
],
|
||||
"txt_in.weight": [
|
||||
"context_embedder.weight",
|
||||
],
|
||||
"txt_in.bias": [
|
||||
"context_embedder.bias",
|
||||
],
|
||||
"img_in.weight": [
|
||||
"x_embedder.weight",
|
||||
],
|
||||
"img_in.bias": [
|
||||
"x_embedder.bias",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.weight": [
|
||||
"norm1.linear.weight",
|
||||
],
|
||||
"double_blocks.().img_mod.lin.bias": [
|
||||
"norm1.linear.bias",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.weight": [
|
||||
"norm1_context.linear.weight",
|
||||
],
|
||||
"double_blocks.().txt_mod.lin.bias": [
|
||||
"norm1_context.linear.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.qkv.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.weight": [
|
||||
"attn.add_q_proj.weight",
|
||||
"attn.add_k_proj.weight",
|
||||
"attn.add_v_proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.qkv.bias": [
|
||||
"attn.add_q_proj.bias",
|
||||
"attn.add_k_proj.bias",
|
||||
"attn.add_v_proj.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.query_norm.scale": [
|
||||
"attn.norm_added_q.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.norm.key_norm.scale": [
|
||||
"attn.norm_added_k.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.weight": [
|
||||
"ff.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.0.bias": [
|
||||
"ff.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.weight": [
|
||||
"ff.net.2.weight",
|
||||
],
|
||||
"double_blocks.().img_mlp.2.bias": [
|
||||
"ff.net.2.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.weight": [
|
||||
"ff_context.net.0.proj.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.0.bias": [
|
||||
"ff_context.net.0.proj.bias",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.weight": [
|
||||
"ff_context.net.2.weight",
|
||||
],
|
||||
"double_blocks.().txt_mlp.2.bias": [
|
||||
"ff_context.net.2.bias",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.weight": [
|
||||
"attn.to_out.0.weight",
|
||||
],
|
||||
"double_blocks.().img_attn.proj.bias": [
|
||||
"attn.to_out.0.bias",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.weight": [
|
||||
"attn.to_add_out.weight",
|
||||
],
|
||||
"double_blocks.().txt_attn.proj.bias": [
|
||||
"attn.to_add_out.bias",
|
||||
],
|
||||
"single_blocks.().modulation.lin.weight": [
|
||||
"norm.linear.weight",
|
||||
],
|
||||
"single_blocks.().modulation.lin.bias": [
|
||||
"norm.linear.bias",
|
||||
],
|
||||
"single_blocks.().linear1.weight": [
|
||||
"attn.to_q.weight",
|
||||
"attn.to_k.weight",
|
||||
"attn.to_v.weight",
|
||||
"proj_mlp.weight",
|
||||
],
|
||||
"single_blocks.().linear1.bias": [
|
||||
"attn.to_q.bias",
|
||||
"attn.to_k.bias",
|
||||
"attn.to_v.bias",
|
||||
"proj_mlp.bias",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().norm.query_norm.scale": [
|
||||
"attn.norm_q.weight",
|
||||
],
|
||||
"single_blocks.().norm.key_norm.scale": [
|
||||
"attn.norm_k.weight",
|
||||
],
|
||||
"single_blocks.().linear2.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"single_blocks.().linear2.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.linear.weight": [
|
||||
"proj_out.weight",
|
||||
],
|
||||
"final_layer.linear.bias": [
|
||||
"proj_out.bias",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.weight": [
|
||||
"norm_out.linear.weight",
|
||||
],
|
||||
"final_layer.adaLN_modulation.1.bias": [
|
||||
"norm_out.linear.bias",
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def is_in_diffusers_map(k):
|
||||
for values in diffusers_map.values():
|
||||
for value in values:
|
||||
if k.endswith(value):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
diffusers = {k: Path.joinpath(diffusers_path, v)
|
||||
for k, v in original_json["weight_map"].items() if is_in_diffusers_map(k)}
|
||||
|
||||
original_safetensors = set(diffusers.values())
|
||||
|
||||
# determine the number of transformer blocks
|
||||
transformer_blocks = 0
|
||||
single_transformer_blocks = 0
|
||||
for key in diffusers.keys():
|
||||
print(key)
|
||||
if key.startswith("transformer_blocks."):
|
||||
print(key)
|
||||
block = int(key.split(".")[1])
|
||||
if block >= transformer_blocks:
|
||||
transformer_blocks = block + 1
|
||||
elif key.startswith("single_transformer_blocks."):
|
||||
block = int(key.split(".")[1])
|
||||
if block >= single_transformer_blocks:
|
||||
single_transformer_blocks = block + 1
|
||||
|
||||
print(f"Transformer blocks: {transformer_blocks}")
|
||||
print(f"Single transformer blocks: {single_transformer_blocks}")
|
||||
|
||||
for file in original_safetensors:
|
||||
if not file.exists():
|
||||
print(f"Error: Missing transformer safetensors file: {file}")
|
||||
exit()
|
||||
|
||||
original_safetensors = {f: safetensors.safe_open(
|
||||
f, framework="pt", device="cpu") for f in original_safetensors}
|
||||
|
||||
|
||||
def swap_scale_shift(weight):
|
||||
shift, scale = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([scale, shift], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
flux_values = {}
|
||||
|
||||
for b in range(transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("double_blocks."):
|
||||
block_prefix = f"transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
for b in range(single_transformer_blocks):
|
||||
for key, weights in diffusers_map.items():
|
||||
if key.startswith("single_blocks."):
|
||||
block_prefix = f"single_transformer_blocks.{b}."
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{block_prefix}{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key.replace("()", f"{b}")] = [
|
||||
f"{block_prefix}{weight}" for weight in weights]
|
||||
|
||||
for key, weights in diffusers_map.items():
|
||||
if not (key.startswith("double_blocks.") or key.startswith("single_blocks.")):
|
||||
found = True
|
||||
for weight in weights:
|
||||
if not (f"{weight}" in diffusers):
|
||||
found = False
|
||||
if found:
|
||||
flux_values[key] = [f"{weight}" for weight in weights]
|
||||
|
||||
flux = {}
|
||||
|
||||
for key, values in tqdm.tqdm(flux_values.items()):
|
||||
if len(values) == 1:
|
||||
flux[key] = original_safetensors[diffusers[values[0]]
|
||||
].get_tensor(values[0]).to("cpu")
|
||||
else:
|
||||
flux[key] = torch.cat(
|
||||
[
|
||||
original_safetensors[diffusers[value]
|
||||
].get_tensor(value).to("cpu")
|
||||
for value in values
|
||||
]
|
||||
)
|
||||
|
||||
if "norm_out.linear.weight" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.weight"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.weight"]].get_tensor(
|
||||
"norm_out.linear.weight").to("cpu")
|
||||
)
|
||||
if "norm_out.linear.bias" in diffusers:
|
||||
flux["final_layer.adaLN_modulation.1.bias"] = swap_scale_shift(
|
||||
original_safetensors[diffusers["norm_out.linear.bias"]].get_tensor(
|
||||
"norm_out.linear.bias").to("cpu")
|
||||
)
|
||||
|
||||
|
||||
def stochastic_round_to(tensor, dtype=torch.float8_e4m3fn):
|
||||
# Define the float8 range
|
||||
min_val = torch.finfo(dtype).min
|
||||
max_val = torch.finfo(dtype).max
|
||||
|
||||
# Clip values to float8 range
|
||||
tensor = torch.clamp(tensor, min_val, max_val)
|
||||
|
||||
# Convert to float32 for calculations
|
||||
tensor = tensor.float()
|
||||
|
||||
# Get the nearest representable float8 values
|
||||
lower = torch.floor(tensor * 256) / 256
|
||||
upper = torch.ceil(tensor * 256) / 256
|
||||
|
||||
# Calculate the probability of rounding up
|
||||
prob = (tensor - lower) / (upper - lower)
|
||||
|
||||
# Generate random values for stochastic rounding
|
||||
rand = torch.rand_like(tensor)
|
||||
|
||||
# Perform stochastic rounding
|
||||
rounded = torch.where(rand < prob, upper, lower)
|
||||
|
||||
# Convert back to float8
|
||||
return rounded.to(dtype)
|
||||
|
||||
|
||||
# set all the keys to bf16
|
||||
for key in flux.keys():
|
||||
if do_8_bit:
|
||||
flux[key] = stochastic_round_to(
|
||||
flux[key], torch.float8_e4m3fn).to('cpu')
|
||||
else:
|
||||
flux[key] = flux[key].clone().to('cpu', torch.bfloat16)
|
||||
|
||||
# load the quantized state dict
|
||||
quantized_state_dict = safetensors.torch.load_file(quantized_state_dict_path)
|
||||
|
||||
transformer_pre = "model.diffusion_model."
|
||||
did_print = False
|
||||
# remove old parts
|
||||
for key in list(quantized_state_dict.keys()):
|
||||
if key.startswith(transformer_pre):
|
||||
if not did_print:
|
||||
# print("dtype: ", quantized_state_dict[key].dtype)
|
||||
did_print = True
|
||||
del quantized_state_dict[key]
|
||||
|
||||
# add the new parts
|
||||
for key, value in flux.items():
|
||||
quantized_state_dict[transformer_pre + key] = value
|
||||
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
# date format like 2024-08-01 YYYY-MM-DD
|
||||
meta['modelspec.date'] = date.today().strftime("%Y-%m-%d")
|
||||
meta['modelspec.title'] = "Flex.1-alpha"
|
||||
meta['modelspec.author'] = "Ostris, LLC"
|
||||
meta['modelspec.license'] = "Apache-2.0"
|
||||
meta['modelspec.implementation'] = "https://github.com/black-forest-labs/flux"
|
||||
meta['modelspec.architecture'] = "Flex.1-alpha"
|
||||
meta['modelspec.description'] = "Flex.1-alpha"
|
||||
|
||||
|
||||
os.makedirs(os.path.dirname(flux_path), exist_ok=True)
|
||||
|
||||
print(f"Saving to {flux_path}")
|
||||
|
||||
safetensors.torch.save_file(quantized_state_dict, flux_path, metadata=meta)
|
||||
|
||||
print("Done.")
|
||||
245
scripts/extract_lora_from_flex.py
Normal file
245
scripts/extract_lora_from_flex.py
Normal file
@@ -0,0 +1,245 @@
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
import argparse
|
||||
from collections import OrderedDict
|
||||
|
||||
parser = argparse.ArgumentParser(description="Extract LoRA from Flex")
|
||||
parser.add_argument("--base", type=str, default="ostris/Flex.1-alpha", help="Base model path")
|
||||
parser.add_argument("--tuned", type=str, required=True, help="Tuned model path")
|
||||
parser.add_argument("--output", type=str, required=True, help="Output path for lora")
|
||||
parser.add_argument("--rank", type=int, default=32, help="LoRA rank for extraction")
|
||||
parser.add_argument("--gpu", type=int, default=0, help="GPU to process extraction")
|
||||
parser.add_argument("--full", action="store_true", help="Do a full transformer extraction, not just transformer blocks")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
if True:
|
||||
# set cuda environment variable
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from lycoris.utils import extract_linear, extract_conv, make_sparse
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
base = args.base
|
||||
tuned = args.tuned
|
||||
output_path = args.output
|
||||
dim = args.rank
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
|
||||
state_dict_base = {}
|
||||
state_dict_tuned = {}
|
||||
|
||||
output_dict = {}
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_diff(
|
||||
base_unet,
|
||||
db_unet,
|
||||
mode="fixed",
|
||||
linear_mode_param=0,
|
||||
conv_mode_param=0,
|
||||
extract_device="cpu",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
# small_conv=True,
|
||||
small_conv=False,
|
||||
):
|
||||
UNET_TARGET_REPLACE_MODULE = [
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
]
|
||||
LORA_PREFIX_UNET = "transformer"
|
||||
|
||||
def make_state_dict(
|
||||
prefix,
|
||||
root_module: torch.nn.Module,
|
||||
target_module: torch.nn.Module,
|
||||
target_replace_modules,
|
||||
):
|
||||
loras = {}
|
||||
temp = {}
|
||||
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
temp[name] = module
|
||||
|
||||
for name, module in tqdm(
|
||||
list((n, m) for n, m in target_module.named_modules() if n in temp)
|
||||
):
|
||||
weights = temp[name]
|
||||
lora_name = prefix + "." + name
|
||||
# lora_name = lora_name.replace(".", "_")
|
||||
layer = module.__class__.__name__
|
||||
if 'transformer_blocks' not in lora_name and not args.full:
|
||||
continue
|
||||
|
||||
if layer in {
|
||||
"Linear",
|
||||
"Conv2d",
|
||||
"LayerNorm",
|
||||
"GroupNorm",
|
||||
"GroupNorm32",
|
||||
"Embedding",
|
||||
"LoRACompatibleLinear",
|
||||
"LoRACompatibleConv"
|
||||
}:
|
||||
root_weight = module.weight
|
||||
try:
|
||||
if torch.allclose(root_weight, weights.weight):
|
||||
continue
|
||||
except:
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
module = module.to(extract_device, torch.float32)
|
||||
weights = weights.to(extract_device, torch.float32)
|
||||
|
||||
if mode == "full":
|
||||
decompose_mode = "full"
|
||||
elif layer == "Linear":
|
||||
weight, decompose_mode = extract_linear(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
elif layer == "Conv2d":
|
||||
is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1
|
||||
weight, decompose_mode = extract_conv(
|
||||
(root_weight - weights.weight),
|
||||
mode,
|
||||
linear_mode_param if is_linear else conv_mode_param,
|
||||
device=extract_device,
|
||||
)
|
||||
if decompose_mode == "low rank":
|
||||
extract_a, extract_b, diff = weight
|
||||
if small_conv and not is_linear and decompose_mode == "low rank":
|
||||
dim = extract_a.size(0)
|
||||
(extract_c, extract_a, _), _ = extract_conv(
|
||||
extract_a.transpose(0, 1),
|
||||
"fixed",
|
||||
dim,
|
||||
extract_device,
|
||||
True,
|
||||
)
|
||||
extract_a = extract_a.transpose(0, 1)
|
||||
extract_c = extract_c.transpose(0, 1)
|
||||
loras[f"{lora_name}.lora_mid.weight"] = (
|
||||
extract_c.detach().cpu().contiguous().half()
|
||||
)
|
||||
diff = (
|
||||
(
|
||||
root_weight
|
||||
- torch.einsum(
|
||||
"i j k l, j r, p i -> p r k l",
|
||||
extract_c,
|
||||
extract_a.flatten(1, -1),
|
||||
extract_b.flatten(1, -1),
|
||||
)
|
||||
)
|
||||
.detach()
|
||||
.cpu()
|
||||
.contiguous()
|
||||
)
|
||||
del extract_c
|
||||
else:
|
||||
module = module.to("cpu")
|
||||
weights = weights.to("cpu")
|
||||
continue
|
||||
|
||||
if decompose_mode == "low rank":
|
||||
loras[f"{lora_name}.lora_A.weight"] = (
|
||||
extract_a.detach().cpu().contiguous().half()
|
||||
)
|
||||
loras[f"{lora_name}.lora_B.weight"] = (
|
||||
extract_b.detach().cpu().contiguous().half()
|
||||
)
|
||||
# loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half()
|
||||
if use_bias:
|
||||
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
|
||||
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
|
||||
|
||||
indices = sparse_diff.indices().to(torch.int16)
|
||||
values = sparse_diff.values().half()
|
||||
loras[f"{lora_name}.bias_indices"] = indices
|
||||
loras[f"{lora_name}.bias_values"] = values
|
||||
loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to(
|
||||
torch.int16
|
||||
)
|
||||
del extract_a, extract_b, diff
|
||||
elif decompose_mode == "full":
|
||||
if "Norm" in layer:
|
||||
w_key = "w_norm"
|
||||
b_key = "b_norm"
|
||||
else:
|
||||
w_key = "diff"
|
||||
b_key = "diff_b"
|
||||
weight_diff = module.weight - weights.weight
|
||||
loras[f"{lora_name}.{w_key}"] = (
|
||||
weight_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
if getattr(weights, "bias", None) is not None:
|
||||
bias_diff = module.bias - weights.bias
|
||||
loras[f"{lora_name}.{b_key}"] = (
|
||||
bias_diff.detach().cpu().contiguous().half()
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
module = module.to("cpu", torch.bfloat16)
|
||||
weights = weights.to("cpu", torch.bfloat16)
|
||||
return loras
|
||||
|
||||
all_loras = {}
|
||||
|
||||
all_loras |= make_state_dict(
|
||||
LORA_PREFIX_UNET,
|
||||
base_unet,
|
||||
db_unet,
|
||||
UNET_TARGET_REPLACE_MODULE,
|
||||
)
|
||||
del base_unet, db_unet
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_lora_name = set()
|
||||
for k in all_loras:
|
||||
lora_name, weight = k.rsplit(".", 1)
|
||||
all_lora_name.add(lora_name)
|
||||
print(len(all_lora_name))
|
||||
return all_loras
|
||||
|
||||
|
||||
# find all the .safetensors files and load them
|
||||
print("Loading Base")
|
||||
base_model = FluxTransformer2DModel.from_pretrained(base, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
print("Loading Tuned")
|
||||
tuned_model = FluxTransformer2DModel.from_pretrained(tuned, subfolder="transformer", torch_dtype=torch.bfloat16)
|
||||
|
||||
output_dict = extract_diff(
|
||||
base_model,
|
||||
tuned_model,
|
||||
mode="fixed",
|
||||
linear_mode_param=dim,
|
||||
conv_mode_param=dim,
|
||||
extract_device="cuda",
|
||||
use_bias=False,
|
||||
sparsity=0.98,
|
||||
small_conv=False,
|
||||
)
|
||||
|
||||
meta = OrderedDict()
|
||||
meta['format'] = 'pt'
|
||||
|
||||
save_file(output_dict, output_path, metadata=meta)
|
||||
|
||||
print("Done")
|
||||
3
todo_multigpu.md
Normal file
3
todo_multigpu.md
Normal 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
17
toolkit/accelerator.py
Normal file
@@ -0,0 +1,17 @@
|
||||
from accelerate import Accelerator
|
||||
from diffusers.utils.torch_utils import is_compiled_module
|
||||
|
||||
global_accelerator = None
|
||||
|
||||
|
||||
def get_accelerator() -> Accelerator:
|
||||
global global_accelerator
|
||||
if global_accelerator is None:
|
||||
global_accelerator = Accelerator()
|
||||
return global_accelerator
|
||||
|
||||
def unwrap_model(model):
|
||||
accelerator = get_accelerator()
|
||||
model = accelerator.unwrap_model(model)
|
||||
model = model._orig_mod if is_compiled_module(model) else model
|
||||
return model
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
88
toolkit/dequantize.py
Normal file
@@ -0,0 +1,88 @@
|
||||
|
||||
|
||||
from functools import partial
|
||||
from optimum.quanto.tensor import QTensor
|
||||
import torch
|
||||
|
||||
|
||||
def hacked_state_dict(self, *args, **kwargs):
|
||||
orig_state_dict = self.orig_state_dict(*args, **kwargs)
|
||||
new_state_dict = {}
|
||||
for key, value in orig_state_dict.items():
|
||||
if key.endswith("._scale"):
|
||||
continue
|
||||
if key.endswith(".input_scale"):
|
||||
continue
|
||||
if key.endswith(".output_scale"):
|
||||
continue
|
||||
if key.endswith("._data"):
|
||||
key = key[:-6]
|
||||
scale = orig_state_dict[key + "._scale"]
|
||||
# scale is the original dtype
|
||||
dtype = scale.dtype
|
||||
scale = scale.float()
|
||||
value = value.float()
|
||||
dequantized = value * scale
|
||||
|
||||
# handle input and output scaling if they exist
|
||||
input_scale = orig_state_dict.get(key + ".input_scale")
|
||||
|
||||
if input_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if input_scale.item() != 1.0:
|
||||
raise ValueError("Input scale is not 1.0, cannot dequantize")
|
||||
|
||||
output_scale = orig_state_dict.get(key + ".output_scale")
|
||||
|
||||
if output_scale is not None:
|
||||
# make sure the tensor is 1.0
|
||||
if output_scale.item() != 1.0:
|
||||
raise ValueError("Output scale is not 1.0, cannot dequantize")
|
||||
|
||||
new_state_dict[key] = dequantized.to('cpu', dtype=dtype)
|
||||
else:
|
||||
new_state_dict[key] = value
|
||||
return new_state_dict
|
||||
|
||||
# hacks the state dict so we can dequantize before saving
|
||||
def patch_dequantization_on_save(model):
|
||||
model.orig_state_dict = model.state_dict
|
||||
model.state_dict = partial(hacked_state_dict, model)
|
||||
|
||||
|
||||
def dequantize_parameter(module: torch.nn.Module, param_name: str) -> bool:
|
||||
"""
|
||||
Convert a quantized parameter back to a regular Parameter with floating point values.
|
||||
|
||||
Args:
|
||||
module: The module containing the parameter to unquantize
|
||||
param_name: Name of the parameter to unquantize (e.g., 'weight', 'bias')
|
||||
|
||||
Returns:
|
||||
bool: True if parameter was unquantized, False if it was already unquantized
|
||||
"""
|
||||
|
||||
# Check if the parameter exists
|
||||
if not hasattr(module, param_name):
|
||||
raise AttributeError(f"Module has no parameter named '{param_name}'")
|
||||
|
||||
param = getattr(module, param_name)
|
||||
|
||||
# If it's not a parameter or not quantized, nothing to do
|
||||
if not isinstance(param, torch.nn.Parameter):
|
||||
raise TypeError(f"'{param_name}' is not a Parameter")
|
||||
if not isinstance(param, QTensor):
|
||||
return False
|
||||
|
||||
# Convert to float tensor while preserving device and requires_grad
|
||||
with torch.no_grad():
|
||||
float_tensor = param.float()
|
||||
new_param = torch.nn.Parameter(
|
||||
float_tensor,
|
||||
requires_grad=param.requires_grad
|
||||
)
|
||||
|
||||
# Replace the parameter
|
||||
setattr(module, param_name, new_param)
|
||||
|
||||
return True
|
||||
@@ -5,6 +5,7 @@ from typing import Iterable, Optional
|
||||
import weakref
|
||||
import copy
|
||||
import contextlib
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic
|
||||
|
||||
import torch
|
||||
|
||||
@@ -43,9 +44,10 @@ class ExponentialMovingAverage:
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter] = None,
|
||||
decay: float = 0.995,
|
||||
use_num_updates: bool = True,
|
||||
use_num_updates: bool = False,
|
||||
# feeds back the decat to the parameter
|
||||
use_feedback: bool = False
|
||||
use_feedback: bool = False,
|
||||
param_multiplier: float = 1.0
|
||||
):
|
||||
if parameters is None:
|
||||
raise ValueError("parameters must be provided")
|
||||
@@ -54,6 +56,7 @@ class ExponentialMovingAverage:
|
||||
self.decay = decay
|
||||
self.num_updates = 0 if use_num_updates else None
|
||||
self.use_feedback = use_feedback
|
||||
self.param_multiplier = param_multiplier
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [
|
||||
p.clone().detach()
|
||||
@@ -121,13 +124,32 @@ class ExponentialMovingAverage:
|
||||
one_minus_decay = 1.0 - decay
|
||||
with torch.no_grad():
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
tmp = (s_param - param)
|
||||
s_param_float = s_param.float()
|
||||
if s_param.dtype != torch.float32:
|
||||
s_param_float = s_param_float.to(torch.float32)
|
||||
param_float = param
|
||||
if param.dtype != torch.float32:
|
||||
param_float = param_float.to(torch.float32)
|
||||
tmp = (s_param_float - param_float)
|
||||
# tmp will be a new tensor so we can do in-place
|
||||
tmp.mul_(one_minus_decay)
|
||||
s_param.sub_(tmp)
|
||||
|
||||
s_param_float.sub_(tmp)
|
||||
|
||||
update_param = False
|
||||
if self.use_feedback:
|
||||
param.add_(tmp)
|
||||
param_float.add_(tmp)
|
||||
update_param = True
|
||||
|
||||
if self.param_multiplier != 1.0:
|
||||
param_float.mul_(self.param_multiplier)
|
||||
update_param = True
|
||||
|
||||
if s_param.dtype != torch.float32:
|
||||
copy_stochastic(s_param, s_param_float)
|
||||
|
||||
if update_param and param.dtype != torch.float32:
|
||||
copy_stochastic(param, param_float)
|
||||
|
||||
|
||||
def copy_to(
|
||||
self,
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -269,16 +269,7 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
is_active = self.adapter_ref().is_active
|
||||
input_ndim = hidden_states.ndim
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
context_input_ndim = encoder_hidden_states.ndim
|
||||
if context_input_ndim == 4:
|
||||
batch_size, channel, height, width = encoder_hidden_states.shape
|
||||
encoder_hidden_states = encoder_hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size = encoder_hidden_states.shape[0]
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
@@ -297,7 +288,44 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# will be none if disabled
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `context` projections.
|
||||
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
|
||||
|
||||
# attention
|
||||
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
|
||||
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
|
||||
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# begin ip adapter
|
||||
if not is_active:
|
||||
ip_hidden_states = None
|
||||
else:
|
||||
@@ -309,47 +337,6 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
raise ValueError("Unconditional is None but should not be")
|
||||
ip_hidden_states = torch.cat([self.adapter_ref().last_unconditional, ip_hidden_states], dim=0)
|
||||
|
||||
# `context` projections.
|
||||
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
|
||||
|
||||
# attention
|
||||
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
|
||||
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
|
||||
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
# YiYi to-do: update uising apply_rotary_emb
|
||||
# from ..embeddings import apply_rotary_emb
|
||||
# query = apply_rotary_emb(query, image_rotary_emb)
|
||||
# key = apply_rotary_emb(key, image_rotary_emb)
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# do ip adapter
|
||||
# will be none if disabled
|
||||
if ip_hidden_states is not None:
|
||||
# apply scaler
|
||||
if self.train_scaler:
|
||||
@@ -365,8 +352,6 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||
ip_hidden_states = F.scaled_dot_product_attention(
|
||||
query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
@@ -376,26 +361,23 @@ class CustomIPFluxAttnProcessor2_0(torch.nn.Module):
|
||||
|
||||
scale = self.scale
|
||||
hidden_states = hidden_states + scale * ip_hidden_states
|
||||
# end ip adapter
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
if context_input_ndim == 4:
|
||||
encoder_hidden_states = encoder_hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
|
||||
# loosely based on # ref https://github.com/tencent-ailab/IP-Adapter/blob/main/tutorial_train.py
|
||||
class IPAdapter(torch.nn.Module):
|
||||
@@ -659,9 +641,9 @@ class IPAdapter(torch.nn.Module):
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"transformer_blocks.{i}.attn")
|
||||
|
||||
# single transformer blocks do not have cross attn
|
||||
# for i, module in transformer.single_transformer_blocks.named_children():
|
||||
# attn_processor_keys.append(f"single_transformer_blocks.{i}.attn")
|
||||
# single transformer blocks do not have cross attn, but we will do them anyway
|
||||
for i, module in transformer.single_transformer_blocks.named_children():
|
||||
attn_processor_keys.append(f"single_transformer_blocks.{i}.attn")
|
||||
else:
|
||||
attn_processor_keys = list(sd.unet.attn_processors.keys())
|
||||
|
||||
@@ -695,7 +677,7 @@ class IPAdapter(torch.nn.Module):
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = sd.unet.config['block_out_channels'][block_id]
|
||||
elif name.startswith("transformer"):
|
||||
elif name.startswith("transformer") or name.startswith("single_transformer"):
|
||||
if is_flux:
|
||||
hidden_size = 3072
|
||||
else:
|
||||
@@ -773,11 +755,20 @@ class IPAdapter(torch.nn.Module):
|
||||
transformer: FluxTransformer2DModel = sd.unet
|
||||
for i, module in transformer.transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"transformer_blocks.{i}.attn"]
|
||||
|
||||
# do single blocks too even though they dont have cross attn
|
||||
for i, module in transformer.single_transformer_blocks.named_children():
|
||||
module.attn.processor = attn_procs[f"single_transformer_blocks.{i}.attn"]
|
||||
|
||||
self.adapter_modules = torch.nn.ModuleList(
|
||||
[
|
||||
transformer.transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.transformer_blocks))
|
||||
])
|
||||
] + [
|
||||
transformer.single_transformer_blocks[i].attn.processor for i in
|
||||
range(len(transformer.single_transformer_blocks))
|
||||
]
|
||||
)
|
||||
else:
|
||||
sd.unet.set_attn_processor(attn_procs)
|
||||
self.adapter_modules = torch.nn.ModuleList(sd.unet.attn_processors.values())
|
||||
@@ -1170,13 +1161,13 @@ class IPAdapter(torch.nn.Module):
|
||||
# when training just scaler, we do not train anything else
|
||||
if not self.config.train_scaler:
|
||||
param_groups.append({
|
||||
"params": self.get_non_scaler_parameters(),
|
||||
"params": list(self.get_non_scaler_parameters()),
|
||||
"lr": adapter_lr,
|
||||
})
|
||||
if self.config.train_scaler or self.config.merge_scaler:
|
||||
scaler_lr = adapter_lr if self.config.scaler_lr is None else self.config.scaler_lr
|
||||
param_groups.append({
|
||||
"params": self.get_scaler_parameters(),
|
||||
"params": list(self.get_scaler_parameters()),
|
||||
"lr": scaler_lr,
|
||||
})
|
||||
return param_groups
|
||||
|
||||
84
toolkit/logging.py
Normal file
84
toolkit/logging.py
Normal file
@@ -0,0 +1,84 @@
|
||||
from typing import OrderedDict, Optional
|
||||
from PIL import Image
|
||||
|
||||
from toolkit.config_modules import LoggingConfig
|
||||
|
||||
# Base logger class
|
||||
# This class does nothing, it's just a placeholder
|
||||
class EmptyLogger:
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
# start logging the training
|
||||
def start(self):
|
||||
pass
|
||||
|
||||
# collect the log to send
|
||||
def log(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# send the log
|
||||
def commit(self, step: Optional[int] = None):
|
||||
pass
|
||||
|
||||
# log image
|
||||
def log_image(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
# finish logging
|
||||
def finish(self):
|
||||
pass
|
||||
|
||||
# Wandb logger class
|
||||
# This class logs the data to wandb
|
||||
class WandbLogger(EmptyLogger):
|
||||
def __init__(self, project: str, run_name: str | None, config: OrderedDict) -> None:
|
||||
self.project = project
|
||||
self.run_name = run_name
|
||||
self.config = config
|
||||
|
||||
def start(self):
|
||||
try:
|
||||
import wandb
|
||||
except ImportError:
|
||||
raise ImportError("Failed to import wandb. Please install wandb by running `pip install wandb`")
|
||||
|
||||
# send the whole config to wandb
|
||||
run = wandb.init(project=self.project, name=self.run_name, config=self.config)
|
||||
self.run = run
|
||||
self._log = wandb.log # log function
|
||||
self._image = wandb.Image # image object
|
||||
|
||||
def log(self, *args, **kwargs):
|
||||
# when commit is False, wandb increments the step,
|
||||
# but we don't want that to happen, so we set commit=False
|
||||
self._log(*args, **kwargs, commit=False)
|
||||
|
||||
def commit(self, step: Optional[int] = None):
|
||||
# after overall one step is done, we commit the log
|
||||
# by log empty object with commit=True
|
||||
self._log({}, step=step, commit=True)
|
||||
|
||||
def log_image(
|
||||
self,
|
||||
image: Image,
|
||||
id, # sample index
|
||||
caption: str | None = None, # positive prompt
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
# create a wandb image object and log it
|
||||
image = self._image(image, caption=caption, *args, **kwargs)
|
||||
self._log({f"sample_{id}": image}, commit=False)
|
||||
|
||||
def finish(self):
|
||||
self.run.finish()
|
||||
|
||||
# create logger based on the logging config
|
||||
def create_logger(logging_config: LoggingConfig, all_config: OrderedDict):
|
||||
if logging_config.use_wandb:
|
||||
project_name = logging_config.project_name
|
||||
run_name = logging_config.run_name
|
||||
return WandbLogger(project=project_name, run_name=run_name, config=all_config)
|
||||
else:
|
||||
return EmptyLogger()
|
||||
@@ -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)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
33
toolkit/models/decorator.py
Normal file
33
toolkit/models/decorator.py
Normal file
@@ -0,0 +1,33 @@
|
||||
import torch
|
||||
|
||||
|
||||
class Decorator(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_tokens: int = 4,
|
||||
token_size: int = 4096,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.weight: torch.nn.Parameter = torch.nn.Parameter(
|
||||
torch.randn(num_tokens, token_size)
|
||||
)
|
||||
# ensure it is float32
|
||||
self.weight.data = self.weight.data.float()
|
||||
|
||||
def forward(self, text_embeds: torch.Tensor, is_unconditional=False) -> torch.Tensor:
|
||||
# make sure the param is float32
|
||||
if self.weight.dtype != text_embeds.dtype:
|
||||
self.weight.data = self.weight.data.float()
|
||||
# expand batch to match text_embeds
|
||||
batch_size = text_embeds.shape[0]
|
||||
decorator_embeds = self.weight.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
if is_unconditional:
|
||||
# zero pad the decorator embeds
|
||||
decorator_embeds = torch.zeros_like(decorator_embeds)
|
||||
|
||||
if decorator_embeds.dtype != text_embeds.dtype:
|
||||
decorator_embeds = decorator_embeds.to(text_embeds.dtype)
|
||||
text_embeds = torch.cat((text_embeds, decorator_embeds), dim=-2)
|
||||
|
||||
return text_embeds
|
||||
356
toolkit/models/diffusion_feature_extraction.py
Normal file
356
toolkit/models/diffusion_feature_extraction.py
Normal 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
176
toolkit/models/flux.py
Normal file
@@ -0,0 +1,176 @@
|
||||
|
||||
# forward that bypasses the guidance embedding so it can be avoided during training.
|
||||
from functools import partial
|
||||
from typing import Optional
|
||||
import torch
|
||||
from diffusers import FluxTransformer2DModel
|
||||
|
||||
|
||||
def guidance_embed_bypass_forward(self, timestep, guidance, pooled_projection):
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
timesteps_emb = self.timestep_embedder(
|
||||
timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D)
|
||||
pooled_projections = self.text_embedder(pooled_projection)
|
||||
conditioning = timesteps_emb + pooled_projections
|
||||
return conditioning
|
||||
|
||||
# bypass the forward function
|
||||
|
||||
|
||||
def bypass_flux_guidance(transformer):
|
||||
if hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
# dont bypass if it doesnt have the guidance embedding
|
||||
if not hasattr(transformer.time_text_embed, 'guidance_embedder'):
|
||||
return
|
||||
transformer.time_text_embed._bfg_orig_forward = transformer.time_text_embed.forward
|
||||
transformer.time_text_embed.forward = partial(
|
||||
guidance_embed_bypass_forward, transformer.time_text_embed
|
||||
)
|
||||
|
||||
# restore the forward function
|
||||
|
||||
|
||||
def restore_flux_guidance(transformer):
|
||||
if not hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
|
||||
return
|
||||
transformer.time_text_embed.forward = transformer.time_text_embed._bfg_orig_forward
|
||||
del transformer.time_text_embed._bfg_orig_forward
|
||||
|
||||
def new_device_to(self: FluxTransformer2DModel, *args, **kwargs):
|
||||
# Store original device if provided in args or kwargs
|
||||
device_in_kwargs = 'device' in kwargs
|
||||
device_in_args = any(isinstance(arg, (str, torch.device)) for arg in args)
|
||||
|
||||
device = None
|
||||
# Remove device from kwargs if present
|
||||
if device_in_kwargs:
|
||||
device = kwargs['device']
|
||||
del kwargs['device']
|
||||
|
||||
# Only filter args if we detected a device argument
|
||||
if device_in_args:
|
||||
args = list(args)
|
||||
for idx, arg in enumerate(args):
|
||||
if isinstance(arg, (str, torch.device)):
|
||||
device = arg
|
||||
del args[idx]
|
||||
|
||||
self.pos_embed = self.pos_embed.to(device, *args, **kwargs)
|
||||
self.time_text_embed = self.time_text_embed.to(device, *args, **kwargs)
|
||||
self.context_embedder = self.context_embedder.to(device, *args, **kwargs)
|
||||
self.x_embedder = self.x_embedder.to(device, *args, **kwargs)
|
||||
for block in self.transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
for block in self.single_transformer_blocks:
|
||||
block.to(block._split_device, *args, **kwargs)
|
||||
|
||||
self.norm_out = self.norm_out.to(device, *args, **kwargs)
|
||||
self.proj_out = self.proj_out.to(device, *args, **kwargs)
|
||||
|
||||
|
||||
|
||||
return self
|
||||
|
||||
|
||||
|
||||
|
||||
def split_gpu_double_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(self._split_device)
|
||||
if encoder_hidden_states.device != self._split_device:
|
||||
encoder_hidden_states = encoder_hidden_states.to(self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
return self._pre_gpu_split_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
|
||||
|
||||
|
||||
def split_gpu_single_block_forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
temb: torch.FloatTensor,
|
||||
image_rotary_emb=None,
|
||||
joint_attention_kwargs=None,
|
||||
**kwargs
|
||||
):
|
||||
if hidden_states.device != self._split_device:
|
||||
hidden_states = hidden_states.to(device=self._split_device)
|
||||
if temb.device != self._split_device:
|
||||
temb = temb.to(device=self._split_device)
|
||||
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
|
||||
# is a tuple of tensors
|
||||
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
|
||||
|
||||
hidden_state_out = self._pre_gpu_split_forward(hidden_states, temb, image_rotary_emb, joint_attention_kwargs, **kwargs)
|
||||
if hasattr(self, "_split_output_device"):
|
||||
return hidden_state_out.to(self._split_output_device)
|
||||
return hidden_state_out
|
||||
|
||||
|
||||
def add_model_gpu_splitter_to_flux(
|
||||
transformer: FluxTransformer2DModel,
|
||||
# ~ 5 billion for all other params
|
||||
other_module_params: Optional[int] = 5e9,
|
||||
# since they are not trainable, multiply by smaller number
|
||||
other_module_param_count_scale: Optional[float] = 0.3
|
||||
):
|
||||
gpu_id_list = [i for i in range(torch.cuda.device_count())]
|
||||
|
||||
# if len(gpu_id_list) > 2:
|
||||
# raise ValueError("Cannot split to more than 2 GPUs currently.")
|
||||
other_module_params *= other_module_param_count_scale
|
||||
|
||||
# since we are not tuning the
|
||||
total_params = sum(p.numel() for p in transformer.parameters()) + other_module_params
|
||||
|
||||
params_per_gpu = total_params / len(gpu_id_list)
|
||||
|
||||
current_gpu_idx = 0
|
||||
# text encoders, vae, and some non block layers will all be on gpu 0
|
||||
current_gpu_params = other_module_params
|
||||
|
||||
for double_block in transformer.transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
double_block._pre_gpu_split_forward = double_block.forward
|
||||
double_block.forward = partial(
|
||||
split_gpu_double_block_forward, double_block)
|
||||
double_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in double_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
for single_block in transformer.single_transformer_blocks:
|
||||
device = torch.device(f"cuda:{current_gpu_idx}")
|
||||
single_block._pre_gpu_split_forward = single_block.forward
|
||||
single_block.forward = partial(
|
||||
split_gpu_single_block_forward, single_block)
|
||||
single_block._split_device = device
|
||||
# add the params to the current gpu
|
||||
current_gpu_params += sum(p.numel() for p in single_block.parameters())
|
||||
# if the current gpu params are greater than the params per gpu, move to next gpu
|
||||
if current_gpu_params > params_per_gpu:
|
||||
current_gpu_idx += 1
|
||||
current_gpu_params = 0
|
||||
if current_gpu_idx >= len(gpu_id_list):
|
||||
current_gpu_idx = gpu_id_list[-1]
|
||||
|
||||
# add output device to last layer
|
||||
transformer.single_transformer_blocks[-1]._split_output_device = torch.device("cuda:0")
|
||||
|
||||
transformer._pre_gpu_split_to = transformer.to
|
||||
transformer.to = partial(new_device_to, transformer)
|
||||
94
toolkit/models/flux_sage_attn.py
Normal file
94
toolkit/models/flux_sage_attn.py
Normal file
@@ -0,0 +1,94 @@
|
||||
from typing import Optional
|
||||
from diffusers.models.attention_processor import Attention
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class FluxSageAttnProcessor2_0:
|
||||
"""Attention processor used typically in processing the SD3-like self-attention projections."""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("FluxAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: torch.FloatTensor = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
from sageattention import sageattn
|
||||
|
||||
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
|
||||
# `sample` projections.
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
|
||||
if encoder_hidden_states is not None:
|
||||
# `context` projections.
|
||||
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
|
||||
batch_size, -1, attn.heads, head_dim
|
||||
).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
|
||||
|
||||
# attention
|
||||
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
|
||||
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
|
||||
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
hidden_states = sageattn(query, key, value, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states, hidden_states = (
|
||||
hidden_states[:, : encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, encoder_hidden_states.shape[1] :],
|
||||
)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
else:
|
||||
return hidden_states
|
||||
@@ -136,6 +136,8 @@ class InstantLoRAMidModule(torch.nn.Module):
|
||||
def down_forward(self, x, *args, **kwargs):
|
||||
# get the embed
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
down_size = math.prod(self.down_shape)
|
||||
down_weight = self.embed[:, :down_size]
|
||||
|
||||
@@ -170,6 +172,8 @@ class InstantLoRAMidModule(torch.nn.Module):
|
||||
|
||||
def up_forward(self, x, *args, **kwargs):
|
||||
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
|
||||
if x.dtype != self.embed.dtype:
|
||||
x = x.to(self.embed.dtype)
|
||||
up_size = math.prod(self.up_shape)
|
||||
up_weight = self.embed[:, -up_size:]
|
||||
|
||||
@@ -211,7 +215,8 @@ class InstantLoRAModule(torch.nn.Module):
|
||||
vision_tokens: int,
|
||||
head_dim: int,
|
||||
num_heads: int, # number of heads in the resampler
|
||||
sd: 'StableDiffusion'
|
||||
sd: 'StableDiffusion',
|
||||
config=None
|
||||
):
|
||||
super(InstantLoRAModule, self).__init__()
|
||||
# self.linear = torch.nn.Linear(2, 1)
|
||||
|
||||
560
toolkit/models/lumina2.py
Normal file
560
toolkit/models/lumina2.py
Normal 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)
|
||||
618
toolkit/models/pixtral_vision.py
Normal file
618
toolkit/models/pixtral_vision.py
Normal file
@@ -0,0 +1,618 @@
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Any, Union, TYPE_CHECKING
|
||||
import os
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from dataclasses import dataclass
|
||||
from huggingface_hub import snapshot_download
|
||||
from safetensors.torch import load_file
|
||||
import json
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from xformers.ops.fmha.attn_bias import BlockDiagonalMask
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def _norm(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, dim: int, hidden_dim: int, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
||||
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# type: ignore
|
||||
return self.w2(nn.functional.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
def repeat_kv(keys: torch.Tensor, values: torch.Tensor, repeats: int, dim: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
keys = torch.repeat_interleave(keys, repeats=repeats, dim=dim)
|
||||
values = torch.repeat_interleave(values, repeats=repeats, dim=dim)
|
||||
return keys, values
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
xq: torch.Tensor,
|
||||
xk: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
|
||||
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
|
||||
freqs_cis = freqs_cis[:, None, :]
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(-2)
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(-2)
|
||||
return xq_out.type_as(xq), xk_out.type_as(xk)
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
head_dim: int,
|
||||
n_kv_heads: int,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.n_heads: int = n_heads
|
||||
self.head_dim: int = head_dim
|
||||
self.n_kv_heads: int = n_kv_heads
|
||||
|
||||
self.repeats = self.n_heads // self.n_kv_heads
|
||||
|
||||
self.scale = self.head_dim ** -0.5
|
||||
|
||||
self.wq = nn.Linear(dim, n_heads * head_dim, bias=False)
|
||||
self.wk = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
|
||||
self.wv = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
|
||||
self.wo = nn.Linear(n_heads * head_dim, dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
cache: Optional[Any] = None,
|
||||
mask: Optional['BlockDiagonalMask'] = None,
|
||||
) -> torch.Tensor:
|
||||
from xformers.ops.fmha import memory_efficient_attention
|
||||
assert mask is None or cache is None
|
||||
seqlen_sum, _ = x.shape
|
||||
|
||||
xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
|
||||
xq = xq.view(seqlen_sum, self.n_heads, self.head_dim)
|
||||
xk = xk.view(seqlen_sum, self.n_kv_heads, self.head_dim)
|
||||
xv = xv.view(seqlen_sum, self.n_kv_heads, self.head_dim)
|
||||
xq, xk = apply_rotary_emb(xq, xk, freqs_cis=freqs_cis)
|
||||
|
||||
if cache is None:
|
||||
key, val = xk, xv
|
||||
elif cache.prefill:
|
||||
key, val = cache.interleave_kv(xk, xv)
|
||||
cache.update(xk, xv)
|
||||
else:
|
||||
cache.update(xk, xv)
|
||||
key, val = cache.key, cache.value
|
||||
key = key.view(seqlen_sum * cache.max_seq_len,
|
||||
self.n_kv_heads, self.head_dim)
|
||||
val = val.view(seqlen_sum * cache.max_seq_len,
|
||||
self.n_kv_heads, self.head_dim)
|
||||
|
||||
# Repeat keys and values to match number of query heads
|
||||
key, val = repeat_kv(key, val, self.repeats, dim=1)
|
||||
|
||||
# xformers requires (B=1, S, H, D)
|
||||
xq, key, val = xq[None, ...], key[None, ...], val[None, ...]
|
||||
output = memory_efficient_attention(
|
||||
xq, key, val, mask if cache is None else cache.mask)
|
||||
output = output.view(seqlen_sum, self.n_heads * self.head_dim)
|
||||
|
||||
assert isinstance(output, torch.Tensor)
|
||||
|
||||
return self.wo(output) # type: ignore
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
norm_eps: float,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.n_heads = n_heads
|
||||
self.dim = dim
|
||||
self.attention = Attention(
|
||||
dim=dim,
|
||||
n_heads=n_heads,
|
||||
head_dim=head_dim,
|
||||
n_kv_heads=n_kv_heads,
|
||||
)
|
||||
self.attention_norm = RMSNorm(dim, eps=norm_eps)
|
||||
self.ffn_norm = RMSNorm(dim, eps=norm_eps)
|
||||
|
||||
self.feed_forward: nn.Module
|
||||
self.feed_forward = FeedForward(dim=dim, hidden_dim=hidden_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
cache: Optional[Any] = None,
|
||||
mask: Optional['BlockDiagonalMask'] = None,
|
||||
) -> torch.Tensor:
|
||||
r = self.attention.forward(self.attention_norm(x), freqs_cis, cache)
|
||||
h = x + r
|
||||
r = self.feed_forward.forward(self.ffn_norm(h))
|
||||
out = h + r
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class VisionEncoderArgs:
|
||||
hidden_size: int
|
||||
num_channels: int
|
||||
image_size: int
|
||||
patch_size: int
|
||||
intermediate_size: int
|
||||
num_hidden_layers: int
|
||||
num_attention_heads: int
|
||||
rope_theta: float = 1e4 # for rope-2D
|
||||
image_token_id: int = 10
|
||||
|
||||
|
||||
def precompute_freqs_cis_2d(
|
||||
dim: int,
|
||||
height: int,
|
||||
width: int,
|
||||
theta: float,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
freqs_cis: 2D complex tensor of shape (height, width, dim // 2) to be indexed by
|
||||
(height, width) position tuples
|
||||
"""
|
||||
# (dim / 2) frequency bases
|
||||
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
|
||||
|
||||
h = torch.arange(height, device=freqs.device)
|
||||
w = torch.arange(width, device=freqs.device)
|
||||
|
||||
freqs_h = torch.outer(h, freqs[::2]).float()
|
||||
freqs_w = torch.outer(w, freqs[1::2]).float()
|
||||
freqs_2d = torch.cat(
|
||||
[
|
||||
freqs_h[:, None, :].repeat(1, width, 1),
|
||||
freqs_w[None, :, :].repeat(height, 1, 1),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
return torch.polar(torch.ones_like(freqs_2d), freqs_2d)
|
||||
|
||||
|
||||
def position_meshgrid(
|
||||
patch_embeds_list: list[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
positions = torch.cat(
|
||||
[
|
||||
torch.stack(
|
||||
torch.meshgrid(
|
||||
torch.arange(p.shape[-2]),
|
||||
torch.arange(p.shape[-1]),
|
||||
indexing="ij",
|
||||
),
|
||||
dim=-1,
|
||||
).reshape(-1, 2)
|
||||
for p in patch_embeds_list
|
||||
]
|
||||
)
|
||||
return positions
|
||||
|
||||
|
||||
class PixtralVisionEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 1024,
|
||||
num_channels: int = 3,
|
||||
image_size: int = 1024,
|
||||
patch_size: int = 16,
|
||||
intermediate_size: int = 4096,
|
||||
num_hidden_layers: int = 24,
|
||||
num_attention_heads: int = 16,
|
||||
rope_theta: float = 1e4, # for rope-2D
|
||||
image_token_id: int = 10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.args = VisionEncoderArgs(
|
||||
hidden_size=hidden_size,
|
||||
num_channels=num_channels,
|
||||
image_size=image_size,
|
||||
patch_size=patch_size,
|
||||
intermediate_size=intermediate_size,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
num_attention_heads=num_attention_heads,
|
||||
rope_theta=rope_theta,
|
||||
image_token_id=image_token_id,
|
||||
)
|
||||
args = self.args
|
||||
self.patch_conv = nn.Conv2d(
|
||||
in_channels=args.num_channels,
|
||||
out_channels=args.hidden_size,
|
||||
kernel_size=args.patch_size,
|
||||
stride=args.patch_size,
|
||||
bias=False,
|
||||
)
|
||||
self.ln_pre = RMSNorm(args.hidden_size, eps=1e-5)
|
||||
self.transformer = VisionTransformerBlocks(args)
|
||||
|
||||
head_dim = self.args.hidden_size // self.args.num_attention_heads
|
||||
assert head_dim % 2 == 0, "ROPE requires even head_dim"
|
||||
self._freqs_cis: Optional[torch.Tensor] = None
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_name_or_path: str) -> 'PixtralVisionEncoder':
|
||||
if os.path.isdir(pretrained_model_name_or_path):
|
||||
model_folder = pretrained_model_name_or_path
|
||||
else:
|
||||
model_folder = snapshot_download(pretrained_model_name_or_path)
|
||||
|
||||
# make sure there is a config
|
||||
if not os.path.exists(os.path.join(model_folder, "config.json")):
|
||||
raise ValueError(f"Could not find config.json in {model_folder}")
|
||||
|
||||
# load config
|
||||
with open(os.path.join(model_folder, "config.json"), "r") as f:
|
||||
config = json.load(f)
|
||||
|
||||
model = cls(**config)
|
||||
|
||||
# see if there is a state_dict
|
||||
if os.path.exists(os.path.join(model_folder, "model.safetensors")):
|
||||
state_dict = load_file(os.path.join(
|
||||
model_folder, "model.safetensors"))
|
||||
model.load_state_dict(state_dict)
|
||||
|
||||
return model
|
||||
|
||||
@property
|
||||
def max_patches_per_side(self) -> int:
|
||||
return self.args.image_size // self.args.patch_size
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return next(self.parameters()).device
|
||||
|
||||
@property
|
||||
def freqs_cis(self) -> torch.Tensor:
|
||||
if self._freqs_cis is None:
|
||||
self._freqs_cis = precompute_freqs_cis_2d(
|
||||
dim=self.args.hidden_size // self.args.num_attention_heads,
|
||||
height=self.max_patches_per_side,
|
||||
width=self.max_patches_per_side,
|
||||
theta=self.args.rope_theta,
|
||||
)
|
||||
|
||||
if self._freqs_cis.device != self.device:
|
||||
self._freqs_cis = self._freqs_cis.to(device=self.device)
|
||||
|
||||
return self._freqs_cis
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images: List[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
from xformers.ops.fmha.attn_bias import BlockDiagonalMask
|
||||
"""
|
||||
Args:
|
||||
images: list of N_img images of variable sizes, each of shape (C, H, W)
|
||||
|
||||
Returns:
|
||||
image_features: tensor of token features for all tokens of all images of
|
||||
shape (N_toks, D)
|
||||
"""
|
||||
assert isinstance(
|
||||
images, list), f"Expected list of images, got {type(images)}"
|
||||
assert all(len(img.shape) == 3 for img in
|
||||
images), f"Expected images with shape (C, H, W), got {[img.shape for img in images]}"
|
||||
# pass images through initial convolution independently
|
||||
patch_embeds_list = [self.patch_conv(
|
||||
img.unsqueeze(0)).squeeze(0) for img in images]
|
||||
|
||||
# flatten to a single sequence
|
||||
patch_embeds = torch.cat([p.flatten(1).permute(1, 0)
|
||||
for p in patch_embeds_list], dim=0)
|
||||
patch_embeds = self.ln_pre(patch_embeds)
|
||||
|
||||
# positional embeddings
|
||||
positions = position_meshgrid(patch_embeds_list).to(self.device)
|
||||
freqs_cis = self.freqs_cis[positions[:, 0], positions[:, 1]]
|
||||
|
||||
# pass through Transformer with a block diagonal mask delimiting images
|
||||
mask = BlockDiagonalMask.from_seqlens(
|
||||
[p.shape[-2] * p.shape[-1] for p in patch_embeds_list],
|
||||
)
|
||||
out = self.transformer(patch_embeds, mask=mask, freqs_cis=freqs_cis)
|
||||
|
||||
# remove batch dimension of the single sequence
|
||||
return out # type: ignore[no-any-return]
|
||||
|
||||
|
||||
class VisionLanguageAdapter(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int):
|
||||
super().__init__()
|
||||
self.w_in = nn.Linear(
|
||||
in_dim,
|
||||
out_dim,
|
||||
bias=True,
|
||||
)
|
||||
self.gelu = nn.GELU()
|
||||
self.w_out = nn.Linear(out_dim, out_dim, bias=True)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# type: ignore[no-any-return]
|
||||
return self.w_out(self.gelu(self.w_in(x)))
|
||||
|
||||
|
||||
class VisionTransformerBlocks(nn.Module):
|
||||
def __init__(self, args: VisionEncoderArgs):
|
||||
super().__init__()
|
||||
self.layers = torch.nn.ModuleList()
|
||||
for _ in range(args.num_hidden_layers):
|
||||
self.layers.append(
|
||||
TransformerBlock(
|
||||
dim=args.hidden_size,
|
||||
hidden_dim=args.intermediate_size,
|
||||
n_heads=args.num_attention_heads,
|
||||
n_kv_heads=args.num_attention_heads,
|
||||
head_dim=args.hidden_size // args.num_attention_heads,
|
||||
norm_eps=1e-5,
|
||||
)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
mask: 'BlockDiagonalMask',
|
||||
freqs_cis: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
for layer in self.layers:
|
||||
x = layer(x, mask=mask, freqs_cis=freqs_cis)
|
||||
return x
|
||||
|
||||
|
||||
DATASET_MEAN = [0.48145466, 0.4578275, 0.40821073] # RGB
|
||||
DATASET_STD = [0.26862954, 0.26130258, 0.27577711] # RGB
|
||||
|
||||
|
||||
def normalize(image: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Normalize a tensor image with mean and standard deviation.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): Image to be normalized, shape (C, H, W), values in [0, 1].
|
||||
mean (torch.Tensor): Mean for each channel.
|
||||
std (torch.Tensor): Standard deviation for each channel.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Normalized image with shape (C, H, W).
|
||||
"""
|
||||
assert image.shape[0] == len(mean) == len(
|
||||
std), f"{image.shape=}, {mean.shape=}, {std.shape=}"
|
||||
|
||||
# Reshape mean and std to (C, 1, 1) for broadcasting
|
||||
mean = mean.view(-1, 1, 1)
|
||||
std = std.view(-1, 1, 1)
|
||||
|
||||
return (image - mean) / std
|
||||
|
||||
|
||||
def transform_image(image: torch.Tensor, new_size: tuple[int, int]) -> torch.Tensor:
|
||||
"""
|
||||
Resize and normalize the input image.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): Input image tensor of shape (C, H, W), values in [0, 1].
|
||||
new_size (tuple[int, int]): Target size (height, width) for resizing.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Resized and normalized image tensor of shape (C, new_H, new_W).
|
||||
"""
|
||||
# Resize the image
|
||||
resized_image = torch.nn.functional.interpolate(
|
||||
image.unsqueeze(0),
|
||||
size=new_size,
|
||||
mode='bicubic',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
|
||||
# Normalize the image
|
||||
normalized_image = normalize(
|
||||
resized_image,
|
||||
torch.tensor(DATASET_MEAN, device=image.device, dtype=image.dtype),
|
||||
torch.tensor(DATASET_STD, device=image.device, dtype=image.dtype)
|
||||
)
|
||||
|
||||
return normalized_image
|
||||
|
||||
|
||||
class PixtralVisionImagePreprocessor:
|
||||
def __init__(self, image_patch_size=16, max_image_size=1024) -> None:
|
||||
self.image_patch_size = image_patch_size
|
||||
self.max_image_size = max_image_size
|
||||
self.image_token = 10
|
||||
|
||||
def _image_to_num_tokens(self, img: torch.Tensor, max_image_size = None) -> Tuple[int, int]:
|
||||
w: Union[int, float]
|
||||
h: Union[int, float]
|
||||
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
|
||||
w, h = img.shape[-1], img.shape[-2]
|
||||
|
||||
# originally, pixtral used the largest of the 2 dimensions, but we
|
||||
# will use the base size of the image based on number of pixels.
|
||||
# ratio = max(h / self.max_image_size, w / self.max_image_size) # original
|
||||
|
||||
base_size = int(math.sqrt(w * h))
|
||||
ratio = base_size / max_image_size
|
||||
if ratio > 1:
|
||||
w = round(w / ratio)
|
||||
h = round(h / ratio)
|
||||
|
||||
width_tokens = (w - 1) // self.image_patch_size + 1
|
||||
height_tokens = (h - 1) // self.image_patch_size + 1
|
||||
|
||||
return width_tokens, height_tokens
|
||||
|
||||
def __call__(self, image: torch.Tensor, max_image_size=None) -> torch.Tensor:
|
||||
"""
|
||||
Converts ImageChunks to numpy image arrays and image token ids
|
||||
|
||||
Args:
|
||||
image torch tensor with values 0-1 and shape of (C, H, W)
|
||||
|
||||
Returns:
|
||||
processed_image: tensor of token features for all tokens of all images of
|
||||
"""
|
||||
# should not have batch
|
||||
if len(image.shape) == 4:
|
||||
raise ValueError(
|
||||
f"Expected image with shape (C, H, W), got {image.shape}")
|
||||
|
||||
if image.min() < 0.0 or image.max() > 1.0:
|
||||
raise ValueError(
|
||||
f"image tensor values must be between 0 and 1. Got min: {image.min()}, max: {image.max()}")
|
||||
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
|
||||
w, h = self._image_to_num_tokens(image, max_image_size=max_image_size)
|
||||
assert w > 0
|
||||
assert h > 0
|
||||
|
||||
new_image_size = (
|
||||
w * self.image_patch_size,
|
||||
h * self.image_patch_size,
|
||||
)
|
||||
|
||||
processed_image = transform_image(image, new_image_size)
|
||||
|
||||
return processed_image
|
||||
|
||||
|
||||
class PixtralVisionImagePreprocessorCompatibleReturn:
|
||||
def __init__(self, pixel_values) -> None:
|
||||
self.pixel_values = pixel_values
|
||||
|
||||
|
||||
# Compatable version with ai toolkit flow
|
||||
class PixtralVisionImagePreprocessorCompatible(PixtralVisionImagePreprocessor):
|
||||
def __init__(self, image_patch_size=16, max_image_size=1024) -> None:
|
||||
super().__init__(
|
||||
image_patch_size=image_patch_size,
|
||||
max_image_size=max_image_size
|
||||
)
|
||||
self.size = {
|
||||
'height': max_image_size,
|
||||
'width': max_image_size
|
||||
}
|
||||
self.max_image_size = max_image_size
|
||||
self.image_mean = DATASET_MEAN
|
||||
self.image_std = DATASET_STD
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
images,
|
||||
return_tensors="pt",
|
||||
do_resize=True,
|
||||
do_rescale=False,
|
||||
max_image_size=None,
|
||||
) -> torch.Tensor:
|
||||
if max_image_size is None:
|
||||
max_image_size = self.max_image_size
|
||||
out_stack = []
|
||||
if len(images.shape) == 3:
|
||||
images = images.unsqueeze(0)
|
||||
for i in range(images.shape[0]):
|
||||
image = images[i]
|
||||
processed_image = super().__call__(image, max_image_size=max_image_size)
|
||||
out_stack.append(processed_image)
|
||||
|
||||
output = torch.stack(out_stack, dim=0)
|
||||
return PixtralVisionImagePreprocessorCompatibleReturn(output)
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatibleReturn:
|
||||
def __init__(self, hidden_states) -> None:
|
||||
self.hidden_states = hidden_states
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatibleConfig:
|
||||
def __init__(self):
|
||||
self.image_size = 1024
|
||||
self.hidden_size = 1024
|
||||
self.patch_size = 16
|
||||
|
||||
|
||||
class PixtralVisionEncoderCompatible(PixtralVisionEncoder):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int = 1024,
|
||||
num_channels: int = 3,
|
||||
image_size: int = 1024,
|
||||
patch_size: int = 16,
|
||||
intermediate_size: int = 4096,
|
||||
num_hidden_layers: int = 24,
|
||||
num_attention_heads: int = 16,
|
||||
rope_theta: float = 1e4, # for rope-2D
|
||||
image_token_id: int = 10,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(
|
||||
hidden_size=hidden_size,
|
||||
num_channels=num_channels,
|
||||
image_size=image_size,
|
||||
patch_size=patch_size,
|
||||
intermediate_size=intermediate_size,
|
||||
num_hidden_layers=num_hidden_layers,
|
||||
num_attention_heads=num_attention_heads,
|
||||
rope_theta=rope_theta,
|
||||
image_token_id=image_token_id,
|
||||
)
|
||||
self.config = PixtralVisionEncoderCompatibleConfig()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
images,
|
||||
output_hidden_states=True,
|
||||
) -> torch.Tensor:
|
||||
out_stack = []
|
||||
if len(images.shape) == 3:
|
||||
images = images.unsqueeze(0)
|
||||
for i in range(images.shape[0]):
|
||||
image = images[i]
|
||||
# must be in an array
|
||||
image_output = super().forward([image])
|
||||
out_stack.append(image_output)
|
||||
|
||||
output = torch.stack(out_stack, dim=0)
|
||||
return PixtralVisionEncoderCompatibleReturn([output])
|
||||
26
toolkit/models/redux.py
Normal file
26
toolkit/models/redux.py
Normal file
@@ -0,0 +1,26 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class ReduxImageEncoder(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
redux_dim: int = 1152,
|
||||
txt_in_features: int = 4096,
|
||||
device=None,
|
||||
dtype=None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.redux_dim = redux_dim
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.redux_up = nn.Linear(redux_dim, txt_in_features * 3, dtype=dtype)
|
||||
self.redux_down = nn.Linear(
|
||||
txt_in_features * 3, txt_in_features, dtype=dtype)
|
||||
|
||||
def forward(self, sigclip_embeds) -> torch.Tensor:
|
||||
x = self.redux_up(sigclip_embeds)
|
||||
x = torch.nn.functional.silu(x)
|
||||
|
||||
projected_x = self.redux_down(x)
|
||||
return projected_x
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
361
toolkit/optimizers/adafactor.py
Normal file
361
toolkit/optimizers/adafactor.py
Normal file
@@ -0,0 +1,361 @@
|
||||
import math
|
||||
from typing import List
|
||||
import torch
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, stochastic_grad_accummulation
|
||||
from optimum.quanto import QBytesTensor
|
||||
import random
|
||||
|
||||
|
||||
class Adafactor(torch.optim.Optimizer):
|
||||
"""
|
||||
Adafactor implementation with stochastic rounding accumulation and stochastic rounding on apply.
|
||||
Modified from transformers Adafactor implementation to support stochastic rounding accumulation and apply.
|
||||
|
||||
AdaFactor pytorch implementation can be used as a drop in replacement for Adam original fairseq code:
|
||||
https://github.com/pytorch/fairseq/blob/master/fairseq/optim/adafactor.py
|
||||
|
||||
Paper: *Adafactor: Adaptive Learning Rates with Sublinear Memory Cost* https://arxiv.org/abs/1804.04235 Note that
|
||||
this optimizer internally adjusts the learning rate depending on the `scale_parameter`, `relative_step` and
|
||||
`warmup_init` options. To use a manual (external) learning rate schedule you should set `scale_parameter=False` and
|
||||
`relative_step=False`.
|
||||
|
||||
Arguments:
|
||||
params (`Iterable[nn.parameter.Parameter]`):
|
||||
Iterable of parameters to optimize or dictionaries defining parameter groups.
|
||||
lr (`float`, *optional*):
|
||||
The external learning rate.
|
||||
eps (`Tuple[float, float]`, *optional*, defaults to `(1e-30, 0.001)`):
|
||||
Regularization constants for square gradient and parameter scale respectively
|
||||
clip_threshold (`float`, *optional*, defaults to 1.0):
|
||||
Threshold of root mean square of final gradient update
|
||||
decay_rate (`float`, *optional*, defaults to -0.8):
|
||||
Coefficient used to compute running averages of square
|
||||
beta1 (`float`, *optional*):
|
||||
Coefficient used for computing running averages of gradient
|
||||
weight_decay (`float`, *optional*, defaults to 0.0):
|
||||
Weight decay (L2 penalty)
|
||||
scale_parameter (`bool`, *optional*, defaults to `True`):
|
||||
If True, learning rate is scaled by root mean square
|
||||
relative_step (`bool`, *optional*, defaults to `True`):
|
||||
If True, time-dependent learning rate is computed instead of external learning rate
|
||||
warmup_init (`bool`, *optional*, defaults to `False`):
|
||||
Time-dependent learning rate computation depends on whether warm-up initialization is being used
|
||||
|
||||
This implementation handles low-precision (FP16, bfloat) values, but we have not thoroughly tested.
|
||||
|
||||
Recommended T5 finetuning settings (https://discuss.huggingface.co/t/t5-finetuning-tips/684/3):
|
||||
|
||||
- Training without LR warmup or clip_threshold is not recommended.
|
||||
|
||||
- use scheduled LR warm-up to fixed LR
|
||||
- use clip_threshold=1.0 (https://arxiv.org/abs/1804.04235)
|
||||
- Disable relative updates
|
||||
- Use scale_parameter=False
|
||||
- Additional optimizer operations like gradient clipping should not be used alongside Adafactor
|
||||
|
||||
Example:
|
||||
|
||||
```python
|
||||
Adafactor(model.parameters(), scale_parameter=False, relative_step=False, warmup_init=False, lr=1e-3)
|
||||
```
|
||||
|
||||
Others reported the following combination to work well:
|
||||
|
||||
```python
|
||||
Adafactor(model.parameters(), scale_parameter=True, relative_step=True, warmup_init=True, lr=None)
|
||||
```
|
||||
|
||||
When using `lr=None` with [`Trainer`] you will most likely need to use [`~optimization.AdafactorSchedule`]
|
||||
scheduler as following:
|
||||
|
||||
```python
|
||||
from transformers.optimization import Adafactor, AdafactorSchedule
|
||||
|
||||
optimizer = Adafactor(model.parameters(), scale_parameter=True, relative_step=True, warmup_init=True, lr=None)
|
||||
lr_scheduler = AdafactorSchedule(optimizer)
|
||||
trainer = Trainer(..., optimizers=(optimizer, lr_scheduler))
|
||||
```
|
||||
|
||||
Usage:
|
||||
|
||||
```python
|
||||
# replace AdamW with Adafactor
|
||||
optimizer = Adafactor(
|
||||
model.parameters(),
|
||||
lr=1e-3,
|
||||
eps=(1e-30, 1e-3),
|
||||
clip_threshold=1.0,
|
||||
decay_rate=-0.8,
|
||||
beta1=None,
|
||||
weight_decay=0.0,
|
||||
relative_step=False,
|
||||
scale_parameter=False,
|
||||
warmup_init=False,
|
||||
)
|
||||
```"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr=None,
|
||||
eps=(1e-30, 1e-3),
|
||||
clip_threshold=1.0,
|
||||
decay_rate=-0.8,
|
||||
beta1=None,
|
||||
weight_decay=0.0,
|
||||
scale_parameter=True,
|
||||
relative_step=True,
|
||||
warmup_init=False,
|
||||
do_paramiter_swapping=False,
|
||||
paramiter_swapping_factor=0.1,
|
||||
stochastic_accumulation=True,
|
||||
):
|
||||
if lr is not None and relative_step:
|
||||
raise ValueError(
|
||||
"Cannot combine manual `lr` and `relative_step=True` options")
|
||||
if warmup_init and not relative_step:
|
||||
raise ValueError(
|
||||
"`warmup_init=True` requires `relative_step=True`")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"eps": eps,
|
||||
"clip_threshold": clip_threshold,
|
||||
"decay_rate": decay_rate,
|
||||
"beta1": beta1,
|
||||
"weight_decay": weight_decay,
|
||||
"scale_parameter": scale_parameter,
|
||||
"relative_step": relative_step,
|
||||
"warmup_init": warmup_init,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
|
||||
self.base_lrs: List[float] = [
|
||||
lr for group in self.param_groups
|
||||
]
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# setup stochastic grad accum hooks
|
||||
if stochastic_accumulation:
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
self.do_paramiter_swapping = do_paramiter_swapping
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
self._total_paramiter_size = 0
|
||||
# count total paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
self._total_paramiter_size += torch.numel(param)
|
||||
# pretty print total paramiters with comma seperation
|
||||
print(f"Total training paramiters: {self._total_paramiter_size:,}")
|
||||
|
||||
# needs to be enabled to count paramiters
|
||||
if self.do_paramiter_swapping:
|
||||
self.enable_paramiter_swapping(self.paramiter_swapping_factor)
|
||||
|
||||
|
||||
def enable_paramiter_swapping(self, paramiter_swapping_factor=0.1):
|
||||
self.do_paramiter_swapping = True
|
||||
self.paramiter_swapping_factor = paramiter_swapping_factor
|
||||
# call it an initial time
|
||||
self.swap_paramiters()
|
||||
|
||||
def swap_paramiters(self):
|
||||
all_params = []
|
||||
# deactivate all paramiters
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
param.requires_grad_(False)
|
||||
# remove any grad
|
||||
param.grad = None
|
||||
all_params.append(param)
|
||||
# shuffle all paramiters
|
||||
random.shuffle(all_params)
|
||||
|
||||
# keep activating paramiters until we are going to go over the target paramiters
|
||||
target_paramiters = int(self._total_paramiter_size * self.paramiter_swapping_factor)
|
||||
total_paramiters = 0
|
||||
for param in all_params:
|
||||
total_paramiters += torch.numel(param)
|
||||
if total_paramiters >= target_paramiters:
|
||||
break
|
||||
else:
|
||||
param.requires_grad_(True)
|
||||
|
||||
@staticmethod
|
||||
def _get_lr(param_group, param_state):
|
||||
rel_step_sz = param_group["lr"]
|
||||
if param_group["relative_step"]:
|
||||
min_step = 1e-6 * \
|
||||
param_state["step"] if param_group["warmup_init"] else 1e-2
|
||||
rel_step_sz = min(min_step, 1.0 / math.sqrt(param_state["step"]))
|
||||
param_scale = 1.0
|
||||
if param_group["scale_parameter"]:
|
||||
param_scale = max(param_group["eps"][1], param_state["RMS"])
|
||||
return param_scale * rel_step_sz
|
||||
|
||||
@staticmethod
|
||||
def _get_options(param_group, param_shape):
|
||||
factored = len(param_shape) >= 2
|
||||
use_first_moment = param_group["beta1"] is not None
|
||||
return factored, use_first_moment
|
||||
|
||||
@staticmethod
|
||||
def _rms(tensor):
|
||||
return tensor.norm(2) / (tensor.numel() ** 0.5)
|
||||
|
||||
@staticmethod
|
||||
def _approx_sq_grad(exp_avg_sq_row, exp_avg_sq_col):
|
||||
# copy from fairseq's adafactor implementation:
|
||||
# https://github.com/huggingface/transformers/blob/8395f14de6068012787d83989c3627c3df6a252b/src/transformers/optimization.py#L505
|
||||
r_factor = (exp_avg_sq_row / exp_avg_sq_row.mean(dim=-
|
||||
1, keepdim=True)).rsqrt_().unsqueeze(-1)
|
||||
c_factor = exp_avg_sq_col.unsqueeze(-2).rsqrt()
|
||||
return torch.mul(r_factor, c_factor)
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
# adafactor manages its own lr
|
||||
def get_learning_rates(self):
|
||||
lrs = [
|
||||
self._get_lr(group, self.state[group["params"][0]])
|
||||
for group in self.param_groups
|
||||
if group["params"][0].grad is not None
|
||||
]
|
||||
if len(lrs) == 0:
|
||||
lrs = self.base_lrs # if called before stepping
|
||||
return lrs
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""
|
||||
Performs a single optimization step
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
self.step_hook()
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
for p in group["params"]:
|
||||
if p.grad is None or not p.requires_grad:
|
||||
continue
|
||||
|
||||
grad = p.grad
|
||||
if grad.dtype != torch.float32:
|
||||
grad = grad.to(torch.float32)
|
||||
if grad.is_sparse:
|
||||
raise RuntimeError(
|
||||
"Adafactor does not support sparse gradients.")
|
||||
|
||||
# if p has atts _scale then it is quantized. We need to divide the grad by the scale
|
||||
# if hasattr(p, "_scale"):
|
||||
# grad = grad / p._scale
|
||||
|
||||
state = self.state[p]
|
||||
grad_shape = grad.shape
|
||||
|
||||
factored, use_first_moment = self._get_options(
|
||||
group, grad_shape)
|
||||
# State Initialization
|
||||
if len(state) == 0:
|
||||
state["step"] = 0
|
||||
|
||||
if use_first_moment:
|
||||
# Exponential moving average of gradient values
|
||||
state["exp_avg"] = torch.zeros_like(grad)
|
||||
if factored:
|
||||
state["exp_avg_sq_row"] = torch.zeros(
|
||||
grad_shape[:-1]).to(grad)
|
||||
state["exp_avg_sq_col"] = torch.zeros(
|
||||
grad_shape[:-2] + grad_shape[-1:]).to(grad)
|
||||
else:
|
||||
state["exp_avg_sq"] = torch.zeros_like(grad)
|
||||
|
||||
state["RMS"] = 0
|
||||
else:
|
||||
if use_first_moment:
|
||||
state["exp_avg"] = state["exp_avg"].to(grad)
|
||||
if factored:
|
||||
state["exp_avg_sq_row"] = state["exp_avg_sq_row"].to(
|
||||
grad)
|
||||
state["exp_avg_sq_col"] = state["exp_avg_sq_col"].to(
|
||||
grad)
|
||||
else:
|
||||
state["exp_avg_sq"] = state["exp_avg_sq"].to(grad)
|
||||
|
||||
p_data_fp32 = p
|
||||
|
||||
if isinstance(p_data_fp32, QBytesTensor):
|
||||
p_data_fp32 = p_data_fp32.dequantize()
|
||||
if p.dtype != torch.float32:
|
||||
p_data_fp32 = p_data_fp32.clone().float()
|
||||
|
||||
state["step"] += 1
|
||||
state["RMS"] = self._rms(p_data_fp32)
|
||||
lr = self._get_lr(group, state)
|
||||
|
||||
beta2t = 1.0 - math.pow(state["step"], group["decay_rate"])
|
||||
eps = group["eps"]
|
||||
if isinstance(eps, tuple) or isinstance(eps, list):
|
||||
eps = eps[0]
|
||||
update = (grad**2) + eps
|
||||
if factored:
|
||||
exp_avg_sq_row = state["exp_avg_sq_row"]
|
||||
exp_avg_sq_col = state["exp_avg_sq_col"]
|
||||
|
||||
exp_avg_sq_row.mul_(beta2t).add_(
|
||||
update.mean(dim=-1), alpha=(1.0 - beta2t))
|
||||
exp_avg_sq_col.mul_(beta2t).add_(
|
||||
update.mean(dim=-2), alpha=(1.0 - beta2t))
|
||||
|
||||
# Approximation of exponential moving average of square of gradient
|
||||
update = self._approx_sq_grad(
|
||||
exp_avg_sq_row, exp_avg_sq_col)
|
||||
update.mul_(grad)
|
||||
else:
|
||||
exp_avg_sq = state["exp_avg_sq"]
|
||||
|
||||
exp_avg_sq.mul_(beta2t).add_(update, alpha=(1.0 - beta2t))
|
||||
update = exp_avg_sq.rsqrt().mul_(grad)
|
||||
|
||||
update.div_(
|
||||
(self._rms(update) / group["clip_threshold"]).clamp_(min=1.0))
|
||||
update.mul_(lr)
|
||||
|
||||
if use_first_moment:
|
||||
exp_avg = state["exp_avg"]
|
||||
exp_avg.mul_(group["beta1"]).add_(
|
||||
update, alpha=(1 - group["beta1"]))
|
||||
update = exp_avg
|
||||
|
||||
if group["weight_decay"] != 0:
|
||||
p_data_fp32.add_(
|
||||
p_data_fp32, alpha=(-group["weight_decay"] * lr))
|
||||
|
||||
p_data_fp32.add_(-update)
|
||||
|
||||
if p.dtype != torch.float32:
|
||||
# apply stochastic rounding
|
||||
copy_stochastic(p, p_data_fp32)
|
||||
|
||||
return loss
|
||||
162
toolkit/optimizers/adam8bit.py
Normal file
162
toolkit/optimizers/adam8bit.py
Normal file
@@ -0,0 +1,162 @@
|
||||
import math
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, Auto8bitTensor, stochastic_grad_accummulation
|
||||
|
||||
class Adam8bit(Optimizer):
|
||||
"""
|
||||
Implements Adam optimizer with 8-bit state storage and stochastic rounding.
|
||||
|
||||
Arguments:
|
||||
params (iterable): Iterable of parameters to optimize or dicts defining parameter groups
|
||||
lr (float): Learning rate (default: 1e-3)
|
||||
betas (tuple): Coefficients for computing running averages of gradient and its square (default: (0.9, 0.999))
|
||||
eps (float): Term added to denominator to improve numerical stability (default: 1e-8)
|
||||
weight_decay (float): Weight decay coefficient (default: 0)
|
||||
decouple (bool): Use AdamW style decoupled weight decay (default: True)
|
||||
"""
|
||||
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8,
|
||||
weight_decay=0, decouple=True):
|
||||
if not 0.0 <= lr:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if not 0.0 <= eps:
|
||||
raise ValueError(f"Invalid epsilon value: {eps}")
|
||||
if not 0.0 <= betas[0] < 1.0:
|
||||
raise ValueError(f"Invalid beta parameter at index 0: {betas[0]}")
|
||||
if not 0.0 <= betas[1] < 1.0:
|
||||
raise ValueError(f"Invalid beta parameter at index 1: {betas[1]}")
|
||||
|
||||
defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay,
|
||||
decouple=decouple)
|
||||
super(Adam8bit, self).__init__(params, defaults)
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# Setup stochastic grad accumulation hooks
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_memory_efficient_fp16(self):
|
||||
return False
|
||||
|
||||
@property
|
||||
def supports_flat_params(self):
|
||||
return True
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# Copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""Performs a single optimization step.
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model and returns the loss.
|
||||
"""
|
||||
# Call pre step
|
||||
self.step_hook()
|
||||
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
beta1, beta2 = group['betas']
|
||||
eps = group['eps']
|
||||
lr = group['lr']
|
||||
decay = group['weight_decay']
|
||||
decouple = group['decouple']
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
# Apply weight decay (coupled variant)
|
||||
if decay != 0 and not decouple:
|
||||
grad.add_(p_fp32.data, alpha=decay)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
# State initialization
|
||||
if len(state) == 0:
|
||||
state['step'] = 0
|
||||
# Exponential moving average of gradient values
|
||||
state['exp_avg'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
# Exponential moving average of squared gradient values
|
||||
state['exp_avg_sq'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
state['step'] += 1
|
||||
bias_correction1 = 1 - beta1 ** state['step']
|
||||
bias_correction2 = 1 - beta2 ** state['step']
|
||||
|
||||
# Adam EMA updates
|
||||
exp_avg.mul_(beta1).add_(grad, alpha=1-beta1)
|
||||
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1-beta2)
|
||||
|
||||
# Apply weight decay (decoupled variant)
|
||||
if decay != 0 and decouple:
|
||||
p_fp32.data.mul_(1 - lr * decay)
|
||||
|
||||
# Bias correction
|
||||
step_size = lr / bias_correction1
|
||||
denom = (exp_avg_sq.sqrt() / math.sqrt(bias_correction2)).add_(eps)
|
||||
|
||||
# Take step
|
||||
p_fp32.data.addcdiv_(exp_avg, denom, value=-step_size)
|
||||
|
||||
# Update state with stochastic rounding
|
||||
state['exp_avg'] = Auto8bitTensor(exp_avg)
|
||||
state['exp_avg_sq'] = Auto8bitTensor(exp_avg_sq)
|
||||
|
||||
# Apply stochastic rounding to parameters
|
||||
copy_stochastic(p.data, p_fp32.data)
|
||||
|
||||
return loss
|
||||
|
||||
def state_dict(self):
|
||||
"""Returns the state of the optimizer as a dict."""
|
||||
state_dict = super().state_dict()
|
||||
|
||||
# Convert Auto8bitTensor objects to regular state dicts
|
||||
for param_id, param_state in state_dict['state'].items():
|
||||
for key, value in param_state.items():
|
||||
if isinstance(value, Auto8bitTensor):
|
||||
param_state[key] = {
|
||||
'_type': 'Auto8bitTensor',
|
||||
'state': value.state_dict()
|
||||
}
|
||||
|
||||
return state_dict
|
||||
|
||||
def load_state_dict(self, state_dict):
|
||||
"""Loads the optimizer state."""
|
||||
# First, load the basic state
|
||||
super().load_state_dict(state_dict)
|
||||
|
||||
# Then convert any Auto8bitTensor states back to objects
|
||||
for param_id, param_state in self.state.items():
|
||||
for key, value in param_state.items():
|
||||
if isinstance(value, dict) and value.get('_type') == 'Auto8bitTensor':
|
||||
param_state[key] = Auto8bitTensor(value['state'])
|
||||
|
||||
335
toolkit/optimizers/automagic.py
Normal file
335
toolkit/optimizers/automagic.py
Normal 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)
|
||||
256
toolkit/optimizers/optimizer_utils.py
Normal file
256
toolkit/optimizers/optimizer_utils.py
Normal file
@@ -0,0 +1,256 @@
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from typing import Optional
|
||||
from optimum.quanto import QBytesTensor
|
||||
|
||||
|
||||
def compute_scale_for_dtype(tensor, dtype):
|
||||
"""
|
||||
Compute appropriate scale for the given tensor and target dtype.
|
||||
|
||||
Args:
|
||||
tensor: Input tensor to be quantized
|
||||
dtype: Target dtype for quantization
|
||||
Returns:
|
||||
Appropriate scale factor for the quantization
|
||||
"""
|
||||
if dtype == torch.int8:
|
||||
abs_max = torch.max(torch.abs(tensor))
|
||||
return abs_max / 127.0 if abs_max > 0 else 1.0
|
||||
elif dtype == torch.uint8:
|
||||
max_val = torch.max(tensor)
|
||||
min_val = torch.min(tensor)
|
||||
range_val = max_val - min_val
|
||||
return range_val / 255.0 if range_val > 0 else 1.0
|
||||
elif dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
# For float8, we typically want to preserve the magnitude of the values
|
||||
# while fitting within the representable range of the format
|
||||
abs_max = torch.max(torch.abs(tensor))
|
||||
if dtype == torch.float8_e4m3fn:
|
||||
# e4m3fn has range [-448, 448] with no infinities
|
||||
max_representable = 448.0
|
||||
else: # torch.float8_e5m2
|
||||
# e5m2 has range [-57344, 57344] with infinities
|
||||
max_representable = 57344.0
|
||||
|
||||
return abs_max / max_representable if abs_max > 0 else 1.0
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype for quantization: {dtype}")
|
||||
|
||||
def quantize_tensor(tensor, dtype):
|
||||
"""
|
||||
Quantize a floating-point tensor to the target dtype with appropriate scaling.
|
||||
|
||||
Args:
|
||||
tensor: Input tensor (float)
|
||||
dtype: Target dtype for quantization
|
||||
Returns:
|
||||
quantized_data: Quantized tensor
|
||||
scale: Scale factor used
|
||||
"""
|
||||
scale = compute_scale_for_dtype(tensor, dtype)
|
||||
|
||||
if dtype == torch.int8:
|
||||
quantized_data = torch.clamp(torch.round(tensor / scale), -128, 127).to(dtype)
|
||||
elif dtype == torch.uint8:
|
||||
quantized_data = torch.clamp(torch.round(tensor / scale), 0, 255).to(dtype)
|
||||
elif dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
# For float8, we scale and then cast directly to the target type
|
||||
# The casting operation will handle the appropriate rounding
|
||||
scaled_tensor = tensor / scale
|
||||
quantized_data = scaled_tensor.to(dtype)
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype for quantization: {dtype}")
|
||||
|
||||
return quantized_data, scale
|
||||
|
||||
|
||||
def update_parameter(target, result_float):
|
||||
"""
|
||||
Updates a parameter tensor, handling both regular torch.Tensor and QBytesTensor cases
|
||||
with proper rescaling for quantized tensors.
|
||||
|
||||
Args:
|
||||
target: The parameter to update (either torch.Tensor or QBytesTensor)
|
||||
result_float: The new values to assign (torch.Tensor)
|
||||
"""
|
||||
if isinstance(target, QBytesTensor):
|
||||
# Get the target dtype from the existing quantized tensor
|
||||
target_dtype = target._data.dtype
|
||||
|
||||
# Handle device placement
|
||||
device = target._data.device
|
||||
result_float = result_float.to(device)
|
||||
|
||||
# Compute new quantized values and scale
|
||||
quantized_data, new_scale = quantize_tensor(result_float, target_dtype)
|
||||
|
||||
# Update the internal tensors with newly computed values
|
||||
target._data.copy_(quantized_data)
|
||||
target._scale.copy_(new_scale)
|
||||
else:
|
||||
# Regular tensor update
|
||||
target.copy_(result_float)
|
||||
|
||||
|
||||
def get_format_params(dtype: torch.dtype) -> tuple[int, int]:
|
||||
"""
|
||||
Returns (mantissa_bits, total_bits) for each format.
|
||||
mantissa_bits excludes the implicit leading 1.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
return 23, 32
|
||||
elif dtype == torch.bfloat16:
|
||||
return 7, 16
|
||||
elif dtype == torch.float16:
|
||||
return 10, 16
|
||||
elif dtype == torch.float8_e4m3fn:
|
||||
return 3, 8
|
||||
elif dtype == torch.float8_e5m2:
|
||||
return 2, 8
|
||||
elif dtype == torch.int8:
|
||||
return 0, 8 # Int8 doesn't have mantissa bits
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {dtype}")
|
||||
|
||||
|
||||
def copy_stochastic(
|
||||
target: torch.Tensor,
|
||||
source: torch.Tensor,
|
||||
eps: Optional[float] = None
|
||||
) -> None:
|
||||
"""
|
||||
Performs stochastic rounding from source tensor to target tensor.
|
||||
|
||||
Args:
|
||||
target: Destination tensor (determines the target format)
|
||||
source: Source tensor (typically float32)
|
||||
eps: Optional minimum value for stochastic rounding (for numerical stability)
|
||||
"""
|
||||
with torch.no_grad():
|
||||
# If target is float32, just copy directly
|
||||
if target.dtype == torch.float32:
|
||||
target.copy_(source)
|
||||
return
|
||||
|
||||
# Special handling for int8
|
||||
if target.dtype == torch.int8:
|
||||
# Scale the source values to utilize the full int8 range
|
||||
scaled = source * 127.0 # Scale to [-127, 127]
|
||||
|
||||
# Add random noise for stochastic rounding
|
||||
noise = torch.rand_like(scaled) - 0.5
|
||||
rounded = torch.round(scaled + noise)
|
||||
|
||||
# Clamp to int8 range
|
||||
clamped = torch.clamp(rounded, -127, 127)
|
||||
target.copy_(clamped.to(torch.int8))
|
||||
return
|
||||
|
||||
mantissa_bits, _ = get_format_params(target.dtype)
|
||||
|
||||
# Convert source to int32 view
|
||||
source_int = source.view(dtype=torch.int32)
|
||||
|
||||
# Calculate number of bits to round
|
||||
bits_to_round = 23 - mantissa_bits # 23 is float32 mantissa bits
|
||||
|
||||
# Create random integers for stochastic rounding
|
||||
rand = torch.randint_like(
|
||||
source,
|
||||
dtype=torch.int32,
|
||||
low=0,
|
||||
high=(1 << bits_to_round),
|
||||
)
|
||||
|
||||
# Add random values to the bits that will be rounded off
|
||||
result = source_int.clone()
|
||||
result.add_(rand)
|
||||
|
||||
# Mask to keep only the bits we want
|
||||
# Create mask with 1s in positions we want to keep
|
||||
mask = (-1) << bits_to_round
|
||||
result.bitwise_and_(mask)
|
||||
|
||||
# Handle minimum value threshold if specified
|
||||
if eps is not None:
|
||||
eps_int = torch.tensor(
|
||||
eps, dtype=torch.float32).view(dtype=torch.int32)
|
||||
zero_mask = (result.abs() < eps_int)
|
||||
result[zero_mask] = torch.sign(source_int[zero_mask]) * eps_int
|
||||
|
||||
# Convert back to float32 view
|
||||
result_float = result.view(dtype=torch.float32)
|
||||
|
||||
# Special handling for float8 formats
|
||||
if target.dtype == torch.float8_e4m3fn:
|
||||
result_float.clamp_(-448.0, 448.0)
|
||||
elif target.dtype == torch.float8_e5m2:
|
||||
result_float.clamp_(-57344.0, 57344.0)
|
||||
|
||||
# Copy the result to the target tensor
|
||||
update_parameter(target, result_float)
|
||||
# target.copy_(result_float)
|
||||
del result, rand, source_int
|
||||
|
||||
|
||||
class Auto8bitTensor:
|
||||
def __init__(self, data: Tensor, *args, **kwargs):
|
||||
if isinstance(data, dict): # Add constructor from state dict
|
||||
self._load_from_state_dict(data)
|
||||
else:
|
||||
abs_max = data.abs().max().item()
|
||||
scale = abs_max / 127.0 if abs_max > 0 else 1.0
|
||||
|
||||
self.quantized = (data / scale).round().clamp(-127, 127).to(torch.int8)
|
||||
self.scale = scale
|
||||
self.orig_dtype = data.dtype
|
||||
|
||||
def dequantize(self) -> Tensor:
|
||||
return self.quantized.to(dtype=torch.float32) * self.scale
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
# Handle the dtype argument whether it's positional or keyword
|
||||
dtype = None
|
||||
if args and isinstance(args[0], torch.dtype):
|
||||
dtype = args[0]
|
||||
args = args[1:]
|
||||
elif 'dtype' in kwargs:
|
||||
dtype = kwargs['dtype']
|
||||
del kwargs['dtype']
|
||||
|
||||
if dtype is not None:
|
||||
# First dequantize then convert to requested dtype
|
||||
return self.dequantize().to(dtype=dtype, *args, **kwargs)
|
||||
|
||||
# If no dtype specified, just pass through to parent
|
||||
return self.dequantize().to(*args, **kwargs)
|
||||
|
||||
def state_dict(self):
|
||||
"""Returns a dictionary containing the current state of the tensor."""
|
||||
return {
|
||||
'quantized': self.quantized,
|
||||
'scale': self.scale,
|
||||
'orig_dtype': self.orig_dtype
|
||||
}
|
||||
|
||||
def _load_from_state_dict(self, state_dict):
|
||||
"""Loads the tensor state from a state dictionary."""
|
||||
self.quantized = state_dict['quantized']
|
||||
self.scale = state_dict['scale']
|
||||
self.orig_dtype = state_dict['orig_dtype']
|
||||
|
||||
def __str__(self):
|
||||
return f"Auto8bitTensor({self.dequantize()})"
|
||||
|
||||
|
||||
def stochastic_grad_accummulation(param):
|
||||
if hasattr(param, "_accum_grad"):
|
||||
grad_fp32 = param._accum_grad.clone().to(torch.float32)
|
||||
grad_fp32.add_(param.grad.to(torch.float32))
|
||||
copy_stochastic(param._accum_grad, grad_fp32)
|
||||
del grad_fp32
|
||||
del param.grad
|
||||
else:
|
||||
param._accum_grad = param.grad.clone()
|
||||
del param.grad
|
||||
286
toolkit/optimizers/prodigy_8bit.py
Normal file
286
toolkit/optimizers/prodigy_8bit.py
Normal file
@@ -0,0 +1,286 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.optim import Optimizer
|
||||
from toolkit.optimizers.optimizer_utils import copy_stochastic, Auto8bitTensor, stochastic_grad_accummulation
|
||||
|
||||
|
||||
class Prodigy8bit(Optimizer):
|
||||
r"""
|
||||
Implements Adam with Prodigy step-sizes.
|
||||
Handles stochastic rounding for various precisions as well as stochastic gradient accumulation.
|
||||
Stores state in 8bit for memory savings.
|
||||
Leave LR set to 1 unless you encounter instability.
|
||||
|
||||
Arguments:
|
||||
params (iterable):
|
||||
Iterable of parameters to optimize or dicts defining parameter groups.
|
||||
lr (float):
|
||||
Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate.
|
||||
betas (Tuple[float, float], optional): coefficients used for computing
|
||||
running averages of gradient and its square (default: (0.9, 0.999))
|
||||
beta3 (float):
|
||||
coefficients for computing the Prodidy stepsize using running averages.
|
||||
If set to None, uses the value of square root of beta2 (default: None).
|
||||
eps (float):
|
||||
Term added to the denominator outside of the root operation to improve numerical stability. (default: 1e-8).
|
||||
weight_decay (float):
|
||||
Weight decay, i.e. a L2 penalty (default: 0).
|
||||
decouple (boolean):
|
||||
Use AdamW style decoupled weight decay
|
||||
use_bias_correction (boolean):
|
||||
Turn on Adam's bias correction. Off by default.
|
||||
safeguard_warmup (boolean):
|
||||
Remove lr from the denominator of D estimate to avoid issues during warm-up stage. Off by default.
|
||||
d0 (float):
|
||||
Initial D estimate for D-adaptation (default 1e-6). Rarely needs changing.
|
||||
d_coef (float):
|
||||
Coefficient in the expression for the estimate of d (default 1.0).
|
||||
Values such as 0.5 and 2.0 typically work as well.
|
||||
Changing this parameter is the preferred way to tune the method.
|
||||
growth_rate (float):
|
||||
prevent the D estimate from growing faster than this multiplicative rate.
|
||||
Default is inf, for unrestricted. Values like 1.02 give a kind of learning
|
||||
rate warmup effect.
|
||||
fsdp_in_use (bool):
|
||||
If you're using sharded parameters, this should be set to True. The optimizer
|
||||
will attempt to auto-detect this, but if you're using an implementation other
|
||||
than PyTorch's builtin version, the auto-detection won't work.
|
||||
"""
|
||||
|
||||
def __init__(self, params, lr=1.0,
|
||||
betas=(0.9, 0.999), beta3=None,
|
||||
eps=1e-8, weight_decay=0, decouple=True,
|
||||
use_bias_correction=False, safeguard_warmup=False,
|
||||
d0=1e-6, d_coef=1.0, growth_rate=float('inf'),
|
||||
fsdp_in_use=False):
|
||||
if not 0.0 < d0:
|
||||
raise ValueError("Invalid d0 value: {}".format(d0))
|
||||
if not 0.0 < lr:
|
||||
raise ValueError("Invalid learning rate: {}".format(lr))
|
||||
if not 0.0 < eps:
|
||||
raise ValueError("Invalid epsilon value: {}".format(eps))
|
||||
if not 0.0 <= betas[0] < 1.0:
|
||||
raise ValueError(
|
||||
"Invalid beta parameter at index 0: {}".format(betas[0]))
|
||||
if not 0.0 <= betas[1] < 1.0:
|
||||
raise ValueError(
|
||||
"Invalid beta parameter at index 1: {}".format(betas[1]))
|
||||
|
||||
if decouple and weight_decay > 0:
|
||||
print(f"Using decoupled weight decay")
|
||||
|
||||
defaults = dict(lr=lr, betas=betas, beta3=beta3,
|
||||
eps=eps, weight_decay=weight_decay,
|
||||
d=d0, d0=d0, d_max=d0,
|
||||
d_numerator=0.0, d_coef=d_coef,
|
||||
k=0, growth_rate=growth_rate,
|
||||
use_bias_correction=use_bias_correction,
|
||||
decouple=decouple, safeguard_warmup=safeguard_warmup,
|
||||
fsdp_in_use=fsdp_in_use)
|
||||
self.d0 = d0
|
||||
super(Prodigy8bit, self).__init__(params, defaults)
|
||||
|
||||
self.is_stochastic_rounding_accumulation = False
|
||||
|
||||
# setup stochastic grad accum hooks
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and param.dtype != torch.float32:
|
||||
self.is_stochastic_rounding_accumulation = True
|
||||
param.register_post_accumulate_grad_hook(
|
||||
stochastic_grad_accummulation
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_memory_efficient_fp16(self):
|
||||
return False
|
||||
|
||||
@property
|
||||
def supports_flat_params(self):
|
||||
return True
|
||||
|
||||
def step_hook(self):
|
||||
if not self.is_stochastic_rounding_accumulation:
|
||||
return
|
||||
# copy over stochastically rounded grads
|
||||
for group in self.param_groups:
|
||||
for param in group['params']:
|
||||
if param.requires_grad and hasattr(param, "_accum_grad"):
|
||||
param.grad = param._accum_grad
|
||||
del param._accum_grad
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
"""Performs a single optimization step.
|
||||
|
||||
Arguments:
|
||||
closure (callable, optional): A closure that reevaluates the model
|
||||
and returns the loss.
|
||||
"""
|
||||
# call pre step
|
||||
self.step_hook()
|
||||
loss = None
|
||||
if closure is not None:
|
||||
loss = closure()
|
||||
|
||||
d_denom = 0.0
|
||||
|
||||
group = self.param_groups[0]
|
||||
use_bias_correction = group['use_bias_correction']
|
||||
beta1, beta2 = group['betas']
|
||||
beta3 = group['beta3']
|
||||
if beta3 is None:
|
||||
beta3 = math.sqrt(beta2)
|
||||
k = group['k']
|
||||
|
||||
d = group['d']
|
||||
d_max = group['d_max']
|
||||
d_coef = group['d_coef']
|
||||
lr = max(group['lr'] for group in self.param_groups)
|
||||
|
||||
if use_bias_correction:
|
||||
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
|
||||
else:
|
||||
bias_correction = 1
|
||||
|
||||
dlr = d*lr*bias_correction
|
||||
|
||||
growth_rate = group['growth_rate']
|
||||
decouple = group['decouple']
|
||||
fsdp_in_use = group['fsdp_in_use']
|
||||
|
||||
d_numerator = group['d_numerator']
|
||||
d_numerator *= beta3
|
||||
|
||||
for group in self.param_groups:
|
||||
decay = group['weight_decay']
|
||||
k = group['k']
|
||||
eps = group['eps']
|
||||
group_lr = group['lr']
|
||||
d0 = group['d0']
|
||||
safeguard_warmup = group['safeguard_warmup']
|
||||
|
||||
if group_lr not in [lr, 0.0]:
|
||||
raise RuntimeError(
|
||||
f"Setting different lr values in different parameter groups is only supported for values of 0")
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
if hasattr(p, "_fsdp_flattened"):
|
||||
fsdp_in_use = True
|
||||
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
# Apply weight decay (coupled variant)
|
||||
if decay != 0 and not decouple:
|
||||
grad.add_(p_fp32.data, alpha=decay)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
# State initialization
|
||||
if 'step' not in state:
|
||||
state['step'] = 0
|
||||
state['s'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
state['p0'] = Auto8bitTensor(p_fp32.detach().clone())
|
||||
# Exponential moving average of gradient values
|
||||
state['exp_avg'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
# Exponential moving average of squared gradient values
|
||||
state['exp_avg_sq'] = Auto8bitTensor(
|
||||
torch.zeros_like(p_fp32.data).detach())
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
s = state['s'].to(torch.float32)
|
||||
p0 = state['p0'].to(torch.float32)
|
||||
|
||||
if group_lr > 0.0:
|
||||
# we use d / d0 instead of just d to avoid getting values that are too small
|
||||
d_numerator += (d / d0) * dlr * torch.dot(grad.flatten(),
|
||||
(p0.data - p_fp32.data).flatten()).item()
|
||||
|
||||
# Adam EMA updates
|
||||
exp_avg.mul_(beta1).add_(grad, alpha=d * (1-beta1))
|
||||
exp_avg_sq.mul_(beta2).addcmul_(
|
||||
grad, grad, value=d * d * (1-beta2))
|
||||
|
||||
if safeguard_warmup:
|
||||
s.mul_(beta3).add_(grad, alpha=((d / d0) * d))
|
||||
else:
|
||||
s.mul_(beta3).add_(grad, alpha=((d / d0) * dlr))
|
||||
d_denom += s.abs().sum().item()
|
||||
|
||||
# update state with stochastic rounding
|
||||
state['exp_avg'] = Auto8bitTensor(exp_avg)
|
||||
state['exp_avg_sq'] = Auto8bitTensor(exp_avg_sq)
|
||||
state['s'] = Auto8bitTensor(s)
|
||||
state['p0'] = Auto8bitTensor(p0)
|
||||
|
||||
d_hat = d
|
||||
|
||||
# if we have not done any progres, return
|
||||
# if we have any gradients available, will have d_denom > 0 (unless \|g\|=0)
|
||||
if d_denom == 0:
|
||||
return loss
|
||||
|
||||
if lr > 0.0:
|
||||
if fsdp_in_use:
|
||||
dist_tensor = torch.zeros(2).cuda()
|
||||
dist_tensor[0] = d_numerator
|
||||
dist_tensor[1] = d_denom
|
||||
dist.all_reduce(dist_tensor, op=dist.ReduceOp.SUM)
|
||||
global_d_numerator = dist_tensor[0]
|
||||
global_d_denom = dist_tensor[1]
|
||||
else:
|
||||
global_d_numerator = d_numerator
|
||||
global_d_denom = d_denom
|
||||
|
||||
d_hat = d_coef * global_d_numerator / global_d_denom
|
||||
if d == group['d0']:
|
||||
d = max(d, d_hat)
|
||||
d_max = max(d_max, d_hat)
|
||||
d = min(d_max, d * growth_rate)
|
||||
|
||||
for group in self.param_groups:
|
||||
group['d_numerator'] = global_d_numerator
|
||||
group['d_denom'] = global_d_denom
|
||||
group['d'] = d
|
||||
group['d_max'] = d_max
|
||||
group['d_hat'] = d_hat
|
||||
|
||||
decay = group['weight_decay']
|
||||
k = group['k']
|
||||
eps = group['eps']
|
||||
|
||||
for p in group['params']:
|
||||
if p.grad is None:
|
||||
continue
|
||||
grad = p.grad.data.to(torch.float32)
|
||||
p_fp32 = p.clone().to(torch.float32)
|
||||
|
||||
state = self.state[p]
|
||||
|
||||
exp_avg = state['exp_avg'].to(torch.float32)
|
||||
exp_avg_sq = state['exp_avg_sq'].to(torch.float32)
|
||||
|
||||
state['step'] += 1
|
||||
|
||||
denom = exp_avg_sq.sqrt().add_(d * eps)
|
||||
|
||||
# Apply weight decay (decoupled variant)
|
||||
if decay != 0 and decouple:
|
||||
p_fp32.data.add_(p_fp32.data, alpha=-decay * dlr)
|
||||
|
||||
# Take step
|
||||
p_fp32.data.addcdiv_(exp_avg, denom, value=-dlr)
|
||||
# apply stochastic rounding
|
||||
copy_stochastic(p.data, p_fp32.data)
|
||||
|
||||
group['k'] = k + 1
|
||||
|
||||
return loss
|
||||
@@ -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
6
toolkit/print.py
Normal 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)
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user