394 Commits

Author SHA1 Message Date
Jaret Burkett
ce4c5291a0 Added experimental wavelet loss 2025-03-26 18:11:23 -06:00
Jaret Burkett
c101f07834 Version bump 2025-03-26 12:16:01 -06:00
Jaret Burkett
e4526ad4a4 Updates to handle video in a dataset on ui 2025-03-26 12:15:28 -06:00
Jaret Burkett
4595965e06 Added an inpainting mask generator for training inpainting if inpaint mask is not provided 2025-03-25 12:16:10 -06:00
Jaret Burkett
41edc18750 Removed unnessary import 2025-03-25 11:54:42 -06:00
Jaret Burkett
6021a3dbc0 Change inpainting mask to zero out on latents instead of image for inpaint area. 2025-03-24 14:16:52 -06:00
Jaret Burkett
71d7a52146 Fixed issue with python being wrong on docker 2025-03-24 14:16:09 -06:00
Jaret Burkett
45be82d5d6 Handle inpainting training for control_lora adapter 2025-03-24 13:17:47 -06:00
Jaret Burkett
f10937e6da Handle multi control inputs for control lora training 2025-03-23 07:37:08 -06:00
Jaret Burkett
ccb66c748f Update readme install directions 2025-03-22 15:14:04 -06:00
Jaret Burkett
2aca2883e7 Update windows install directions for new version of torch 2025-03-22 15:12:51 -06:00
Jaret Burkett
1ad58c5816 Changed control lora to only have new weights and leave other input weights alone for more flexability of using multiple ones together. 2025-03-22 10:24:52 -06:00
Jaret Burkett
6dea41b9fc Version bump 2025-03-21 11:48:03 -06:00
Jaret Burkett
9a902c067f Show the console log on the job overview in the ui. 2025-03-21 11:45:36 -06:00
Jaret Burkett
0bbc69c135 Fix issue where the job would hang in the ui if it failed to start 2025-03-21 09:46:41 -06:00
Jaret Burkett
6c5eb0cf87 Move pytorch install above cache bust to prevent reinstalling and reuploading it 2025-03-21 06:43:09 -06:00
Jaret Burkett
e3373671b9 Fixed missing dependency 2025-03-21 06:25:30 -06:00
Jaret Burkett
aceb3a0f25 Rework docker 2025-03-20 20:07:27 -06:00
Jaret Burkett
c8049a483d Fix issue with auth token not on build time 2025-03-20 17:27:30 -06:00
Jaret Burkett
f5aa4232fa Added ability to quantize with torchao 2025-03-20 16:28:54 -06:00
Jaret Burkett
3a6b24f4c8 Added a way to secure the UI. Plus various bug fixes and quality of life updates 2025-03-20 08:07:09 -06:00
Jaret Burkett
bbfd6ef0fe Fixed bug that prevented using schnell training adapter 2025-03-19 10:25:12 -06:00
Jaret Burkett
b829983b16 Added ability to load video datasets and train with them 2025-03-19 09:54:26 -06:00
Jaret Burkett
fa187b1208 Added differential masking to targeted_flow_guidance to allow the model learn to clean up the targeted area a little more than unmasked was capable of 2025-03-17 13:25:01 -06:00
Jaret Burkett
5eb627dd9d Add targeted flow guidance training for flow based models 2025-03-17 09:21:23 -06:00
Jaret Burkett
604e76d34d Fix issue with full finetuning wan 2025-03-17 09:17:40 -06:00
Jaret Burkett
6cde96ae5f Adding files I forgot to stage to the last commit 2025-03-15 11:59:41 -06:00
Jaret Burkett
1be613ed06 Added advanced mode yaml editor to the ui 2025-03-15 11:58:59 -06:00
Jaret Burkett
c52421aab7 Allow clip image to not have processor on dataloader for raw img 2025-03-15 08:27:54 -06:00
Jaret Burkett
3812957bc9 Added ability to train control loras. Other important bug fixes thrown in 2025-03-14 18:03:00 -06:00
Jaret Burkett
391329dbdc Fix issue with device placement on te 2025-03-13 20:48:12 -06:00
Jaret Burkett
3b45892b4f Update sponsors 2025-03-13 19:06:59 -06:00
Jaret Burkett
cf4216e6b8 Add support for wan training in ui 2025-03-13 18:54:27 -06:00
Jaret Burkett
31e057d9a3 Fixed issue with device placement in some scenereos when doing low vram on wan 2025-03-13 10:30:27 -06:00
Jaret Burkett
d507b44a7b Update sponsors 2025-03-08 22:28:20 -07:00
Jaret Burkett
242c04a0b8 Fix error with training video models with batch greater than 1 2025-03-08 18:47:27 -07:00
Jaret Burkett
386e68a422 Fixed a bug that changes all samples to webp 2025-03-08 18:02:56 -07:00
Jaret Burkett
850b8da6e5 Added siglip 2 vision encoder for custom adapter 2025-03-09 00:14:44 +00:00
Jaret Burkett
51ad19b568 Add config file examples for training Wan LoRAs on 24GB cards. 2025-03-08 13:56:21 -07:00
Jaret Burkett
e6739f7eb2 Convert wan lora weights on save to be something comfy can handle 2025-03-08 12:55:11 -07:00
Jaret Burkett
7e37918fbc Double tap module casting as it doesent seem to happen every time. 2025-03-07 22:15:24 -07:00
Jaret Burkett
4d88f8f218 Fixed cuda error when not all tensors have been moved to the correct device. 2025-03-07 22:04:35 -07:00
Jaret Burkett
25341c4613 Got wan 14b training to work on 24GB card. 2025-03-07 17:04:10 -07:00
Jaret Burkett
391cf80fea Added training for Wan2.1. Not finalized, wait. 2025-03-07 13:53:44 -07:00
Jaret Burkett
4e3bda7c70 Merge pull request #264 from ostris/cogview4
Added basics for CogView4. Broken as hell though. Dont use.
2025-03-05 14:52:06 -07:00
Jaret Burkett
763128ea42 Note about cogview 2025-03-05 14:46:11 -07:00
Jaret Burkett
4fe33f51c1 Fix issue with picking layers for quantization, adjust layers fo better quantization of cogview4 2025-03-05 13:44:40 -07:00
Jaret Burkett
aa44828c0c WIP more work on cogview4 2025-03-05 09:43:00 -07:00
Jaret Burkett
6f6fb90812 Added cogview4. Loss still needs work. 2025-03-04 18:43:52 -07:00
Jaret Burkett
c57434ad7b Removed wan submodule stuff for now 2025-03-04 00:32:24 -07:00
Jaret Burkett
8bb47d1bfe Merge branch 'main' into wan21 2025-03-04 00:31:57 -07:00
Jaret Burkett
e7dbb20f68 Removed wan submodule for now 2025-03-04 00:29:19 -07:00
Jaret Burkett
c5e0c2bbe2 Fixes to allow for redux assisted training 2025-03-03 16:27:19 -07:00
Jaret Burkett
1f3f45a48d Bugfixes 2025-03-03 08:22:15 -07:00
Jaret Burkett
3c8c84f156 Added supporters to readme and a script to update it 2025-03-02 10:25:27 -07:00
Jaret Burkett
b001d77efb Added LoKr instructions to the readme 2025-03-02 08:55:56 -07:00
Jaret Burkett
7ae31c9ae9 Added LoKr to the ui 2025-03-02 08:49:01 -07:00
Jaret Burkett
b16819f8e7 Added LoKr support 2025-03-02 06:57:50 -07:00
Jaret Burkett
f5e40dfa62 WIP on wan 2025-03-01 16:12:52 -07:00
Jaret Burkett
acc79956aa WIP create new class to add new models more easily 2025-03-01 13:49:02 -07:00
Jaret Burkett
60539c0b0f Allow using prior loss with a custom adapter 2025-03-01 08:01:14 -07:00
Jaret Burkett
dd700f70b3 Avoid loading state dict for automagic for now until I can sort out some issues 2025-02-26 17:03:14 -07:00
Jaret Burkett
d360e76661 fixed issue with dop prompt replacement 2025-02-26 13:35:18 -07:00
Jaret Burkett
6ec23ed226 Fixed issue when doing inverted masked prior with flowmatching algos 2025-02-26 12:12:32 -07:00
Jaret Burkett
f6e16e582a Added Differential Output Preservation Loss to trainer and ui 2025-02-25 20:12:36 -07:00
Jaret Burkett
259ded9602 Fixed issue with trigger word saving in ui 2025-02-24 11:04:24 -07:00
Jaret Burkett
440ba5fb3d Spawn windows in an cmd terminal. Should be working now, but not sure on my system 2025-02-24 08:54:56 -07:00
Jaret Burkett
093f14ac19 UI Bug fixes and initial windows support 2025-02-24 08:15:22 -07:00
Jaret Burkett
f0fbd8bb53 Merge pull request #256 from ostris/ui
Added AI-Toolkit UI
2025-02-23 16:10:49 -07:00
Jaret Burkett
0a981bea2b Fixed typo 2025-02-23 16:07:38 -07:00
Jaret Burkett
1d0e3a4498 Fixed some build issues for now. Added info to the readme 2025-02-23 15:59:17 -07:00
Jaret Burkett
3c7daf49f3 Add HF token to env when spawing via ui 2025-02-23 14:52:12 -07:00
Jaret Burkett
56d8d6bd81 Capture speed from the timer for the ui 2025-02-23 14:38:46 -07:00
Jaret Burkett
3e49337a58 Set step to the last step saved at when exiting 2025-02-23 13:21:22 -07:00
Jaret Burkett
60f848a877 Send more data when loading the model to the ui 2025-02-23 12:49:54 -07:00
Jaret Burkett
b366e46f1c Added more settings to the training config 2025-02-23 12:34:52 -07:00
Jaret Burkett
a280f78c69 Added checkpoint downloader 2025-02-22 16:48:15 -07:00
Jaret Burkett
6e19e7449e Fixed some issues with gpu info refreshing 2025-02-22 14:14:23 -07:00
Jaret Burkett
a6d46ad9ae Cleanup of job page 2025-02-22 13:54:06 -07:00
Jaret Burkett
f3725578dd Cleaned up dashboard 2025-02-22 13:23:26 -07:00
Jaret Burkett
ed99c3c0c8 Moved gpu to its own widget 2025-02-22 12:43:20 -07:00
Jaret Burkett
ed84c19205 Moved the job action bar to a shred component 2025-02-22 12:20:14 -07:00
Jaret Burkett
a7a9c11d9e Fixed add image dropbox 2025-02-22 11:59:21 -07:00
Jaret Burkett
f60698d0ee Fixed some bugs with ui and lock job name to prevent issues with continuing training. 2025-02-22 11:49:36 -07:00
Jaret Burkett
5f094fb17a Added controls to the jobs table 2025-02-22 10:57:53 -07:00
Jaret Burkett
a5227cba7b Switched to a universal table library 2025-02-22 09:59:17 -07:00
Jaret Burkett
77a5e01301 Added proper icon 2025-02-22 09:05:55 -07:00
Jaret Burkett
4ef5a668c0 Make left arrow browsing only hit last image max 2025-02-21 22:04:53 -07:00
Jaret Burkett
f081d14527 Preview samples full screen and use arrow keys to navigate them 2025-02-21 21:52:24 -07:00
Jaret Burkett
710c6de1c9 Samples work in ui now 2025-02-21 20:28:52 -07:00
Jaret Burkett
2b6e66e0cb Mor ui work 2025-02-21 12:40:17 -07:00
Jaret Burkett
ab641e014f Added funding github stuff 2025-02-21 17:13:36 +00:00
Jaret Burkett
ad87f72384 Start, stop, monitor jobs from ui working. 2025-02-21 09:49:28 -07:00
Jaret Burkett
d0214c0df9 Make ui more uniform 2025-02-21 06:18:27 -07:00
Jaret Burkett
adcf884c0f Built out the ui trainer plugin with db comminication 2025-02-21 05:53:35 -07:00
Jaret Burkett
f778d979b5 Saving captions is working 2025-02-20 16:17:00 -07:00
Jaret Burkett
db3ccbba33 Handle image deletion 2025-02-20 15:58:10 -07:00
Jaret Burkett
0d2be18a9b Delete datasets 2025-02-20 14:49:03 -07:00
Jaret Burkett
bbc340e545 Cleanup and add hooks 2025-02-20 13:38:58 -07:00
Jaret Burkett
33fdfd6091 Added beginning or lokr 2025-02-20 12:47:42 -07:00
Jaret Burkett
9f6030620f Dataset uploads working 2025-02-20 12:47:01 -07:00
Jaret Burkett
b5252b5028 More ui work 2025-02-20 11:19:01 -07:00
Jaret Burkett
b0d8fc220d More ui more ui 2025-02-19 20:54:02 -07:00
Jaret Burkett
cef7d9e594 Config ui section is coming along 2025-02-19 07:52:24 -07:00
Jaret Burkett
b13fcc1039 Setup a very basic ui 2025-02-18 10:57:14 -07:00
Jaret Burkett
b32d7e552b Shamelessly beg for money 2025-02-18 05:15:29 -07:00
Jaret Burkett
4af6c5cf30 Work on supporting flex.2 potential arch 2025-02-17 14:10:25 -07:00
Jaret Burkett
1f7784510d WIP Flex 2 pipeline 2025-02-16 14:54:29 -07:00
Jaret Burkett
87e557cf1e Bug fixes and improvements to llmadapter 2025-02-15 07:18:07 -07:00
Jaret Burkett
bd8d7dc081 fixed various issues with llm attention masking. Added block training on the llm adapter. 2025-02-14 11:24:01 -07:00
Jaret Burkett
2be6926398 Added back syustem prompt for llm and remove those tokens from the embeddings 2025-02-14 07:23:37 -07:00
Jaret Burkett
87ac031859 Remove system prompt, shouldnt be necessary fo rhow it works. 2025-02-13 08:42:48 -07:00
Jaret Burkett
7679105d52 Added llm text encoder adapter 2025-02-13 08:28:32 -07:00
Jaret Burkett
2622de1e01 DFE tweaks. Adding support for more llms as text encoders 2025-02-13 04:31:49 -07:00
Jaret Burkett
8450aca10e Fixed missed merge conflice and locked diffusers version 2025-02-12 09:40:02 -07:00
Jaret Burkett
0b8a32def7 merged in lumina2 branch 2025-02-12 09:33:03 -07:00
Jaret Burkett
787bb37e76 Small fixed for DFE, polar guidance, and other things 2025-02-12 09:27:44 -07:00
Jaret Burkett
10aa7e9d5e Fixed some breaking changes with diffusers gradient checkpointing. 2025-02-10 09:35:31 -07:00
Jaret Burkett
ed1deb71c4 Added examples for training lumina2 2025-02-08 16:13:18 -07:00
Jaret Burkett
4de6a825fa Update lumina requirements 2025-02-08 15:16:35 -07:00
Jaret Burkett
9a7266275d Wokr on lumina2 2025-02-08 14:52:39 -07:00
Jaret Burkett
d138f07365 Imitial lumina3 support 2025-02-08 10:59:53 -07:00
Jaret Burkett
c6d8eedb94 Added ability to use consistent noise for each image in a dataset by hashing the path and using that as a seed. 2025-02-08 07:13:48 -07:00
Jaret Burkett
af5e760be1 Merge pull request #249 from ostris/accelerate-multi-gpu
Multi gpu support. Other goodies
2025-02-08 07:11:10 -07:00
Jaret Burkett
ff3d54bb5b Make the mean of the mask multiplier be 1.0 for a more balanced loss. 2025-02-06 10:57:06 +00:00
Jaret Burkett
0e75724b4d Lock version of diffusers 2025-02-05 18:14:59 +00:00
Jaret Burkett
376bb1bf6f Lock torch version due to breaking changes 2025-02-04 23:00:57 +00:00
Jaret Burkett
216ab164ce Experimental features and bug fixes 2025-02-04 13:36:34 -07:00
Jaret Burkett
e6180d1e1d Bug fixes 2025-01-31 13:23:01 -07:00
Jaret Burkett
15a57bc89f Add new version of DFE. Kitchen sink 2025-01-31 11:42:27 -07:00
Jaret Burkett
e5355bf8d5 Added train catch to the blank network 2025-01-30 16:15:45 +00:00
Jaret Burkett
34a1c6947a Added flux_shift as timestep type 2025-01-27 07:35:00 -07:00
Jaret Burkett
2141c6e06c Merge remote-tracking branch 'origin/main' into accelerate-multi-gpu 2025-01-26 11:19:34 -07:00
Jaret Burkett
1188cf1e8a Adjust flux sample sampler to handle some new breaking changes in diffusers. 2025-01-26 18:09:21 +00:00
Jaret Burkett
5e663746b8 Working multi gpu training. Still need a lot of tweaks and testing. 2025-01-25 16:46:20 -07:00
Jaret Burkett
441474e81f Added a flag to lora extraction script to do a full transformer extraction. 2025-01-24 09:34:13 -07:00
Jaret Burkett
a6a690f796 Update full fine tune example to only train transformer blocks. 2025-01-24 09:28:34 -07:00
Jaret Burkett
6191f19e55 Added script to convert diffusers model to ComfyUI variant 2025-01-23 21:30:23 -07:00
Jaret Burkett
bbfba0c188 Added v2 of dfp 2025-01-22 16:32:13 -07:00
Jaret Burkett
e1549ad54d Update dfe model arch 2025-01-22 10:37:23 -07:00
Jaret Burkett
04abe57c76 Added weighing to DFE 2025-01-22 08:50:57 -07:00
Jaret Burkett
89dd041b97 Added ability to pair samples with a closer noise with optimal_noise_pairing_samples 2025-01-21 18:30:10 -07:00
Jaret Burkett
29122b1a54 Added code to handle diffusion feature extraction loss 2025-01-21 14:21:34 -07:00
Jaret Burkett
6a8e3d8610 Added a config file for full finetuning flex. Added a lora extraction script for flex 2025-01-20 10:09:01 -07:00
Jaret Burkett
4c8a9e1b88 Added example config to train Flex 2025-01-18 18:03:20 -07:00
Jaret Burkett
fadb2f3a76 Allow quantizing the te independently on flux. added lognorm_blend timestep schedule 2025-01-18 18:02:31 -07:00
Jaret Burkett
4723f23c0d Added ability to split up flux across gpus (experimental). Changed the way timestep scheduling works to prep for more specific schedules. 2024-12-31 07:06:55 -07:00
Jaret Burkett
8ef07a9c36 Added training for an experimental decoratgor embedding. Allow for turning off guidance embedding on flux (for unreleased model). Various bug fixes and modifications 2024-12-15 08:59:27 -07:00
Jaret Burkett
92ce93140e Adjustments to defaults for automagic 2024-11-29 10:28:06 -07:00
Jaret Burkett
f213996aa5 Fixed saving and displaying for automagic 2024-11-29 08:00:22 -07:00
Jaret Burkett
cbe31eaf0a Initial work on a auto adjusting optimizer 2024-11-29 04:48:58 -07:00
Jaret Burkett
67c2e44edb Added support for training flux redux adapters 2024-11-21 20:01:52 -07:00
Jaret Burkett
96d418bb95 Added support for full finetuning flux with randomized param activation. Examples coming soon 2024-11-21 13:05:32 -07:00
Jaret Burkett
894374b2e9 Various bug fixes and optimizations for quantized training. Added untested custom adam8bit optimizer. Did some work on LoRM (dont use) 2024-11-20 09:16:55 -07:00
Jaret Burkett
6509ba4484 Fix seed generation to make it deterministic so it is consistant from gpu to gpu 2024-11-15 12:11:13 -07:00
Jaret Burkett
025ee3dd3d Added ability for adafactor to fully fine tune quantized model. 2024-10-30 16:38:07 -06:00
Jaret Burkett
58f9d01c2b Added adafactor implementation that handles stochastic rounding of update and accumulation 2024-10-30 05:25:57 -06:00
Jaret Burkett
e72b59a8e9 Added experimental 8bit version of prodigy with stochastic rounding and stochastic gradient accumulation. Still testing. 2024-10-29 14:28:28 -06:00
Jaret Burkett
4aa19b5c1d Only quantize flux T5 is also quantizing model. Load TE from original name and path if fine tuning. 2024-10-29 14:25:31 -06:00
Jaret Burkett
4747716867 Fixed issue with adapters not providing gradients with new grad activator 2024-10-29 14:22:10 -06:00
Jaret Burkett
22cd40d7b9 Improvements for full tuning flux. Added debugging launch config for vscode 2024-10-29 04:54:08 -06:00
Jaret Burkett
3400882a80 Added preliminary support for SD3.5-large lora training 2024-10-22 12:21:36 -06:00
Jaret Burkett
9f94c7b61e Added experimental param multiplier to the ema module 2024-10-22 09:25:52 -06:00
Jaret Burkett
bedb8197a2 Fixed issue with sizes for some images being loaded sideways resulting in squished images. 2024-10-20 11:51:29 -06:00
Jaret Burkett
e3ebd73610 Add a projection layer on vision direct when doing image embeds 2024-10-20 10:48:23 -06:00
Jaret Burkett
dd931757cd Merge branch 'main' of github.com:ostris/ai-toolkit 2024-10-20 07:04:29 -06:00
Jaret Burkett
0640cdf569 Handle errors in loading size database 2024-10-20 07:04:19 -06:00
Jaret Burkett
0b048d0dde Locked version of quanto as it breaks in later versions 2024-10-16 22:41:04 +00:00
Jaret Burkett
473d455f44 Process empty clip image if there is not one for reg images when training a custom adapter 2024-10-15 08:28:04 -06:00
Jaret Burkett
ce759ebd8c Normalize the image embeddings on vd adapter forward 2024-10-12 15:09:48 +00:00
Jaret Burkett
628a7923a3 Remove norm on image embeds on custom adapter 2024-10-12 00:43:18 +00:00
Jaret Burkett
3922981996 Added some additional experimental things to the vision direct encoder 2024-10-10 19:42:26 +00:00
Jaret Burkett
ab22674980 Allow for a default caption file in the folder. Minor bug fixes. 2024-10-10 07:31:33 -06:00
Jaret Burkett
9452929300 Apply a mask to the embeds for SD if using T5 encoder 2024-10-04 10:55:20 -06:00
Jaret Burkett
a800c9d19e Add a method to have an inference only lora 2024-10-04 10:06:53 -06:00
Jaret Burkett
28e6f00790 Fixed bug in returning clip image embed to actually return it 2024-10-03 10:49:09 -06:00
Jaret Burkett
67e0aca750 Added ability to load clip pairs randomly from folder. Other small bug fixes 2024-10-03 10:03:49 -06:00
Jaret Burkett
f05224970f Added Vision Languate Adapter usage for pixtral vd adapter 2024-09-29 19:39:56 -06:00
Jaret Burkett
b4f64de4c2 Quick patch to scope xformer imports until a better solution 2024-09-28 15:36:42 -06:00
Jaret Burkett
2e5f6668dc Add xformers ad a dependency 2024-09-28 15:30:14 -06:00
Jaret Burkett
e4c82803e1 Handle random resizing for pixtral input on direct vision adapter 2024-09-28 14:53:38 -06:00
Jaret Burkett
69aa92bce5 Added support for AdEMAMix8bit 2024-09-28 14:33:51 -06:00
Jaret Burkett
a508caad1d Change pixtral to crop based on number of pixels instead of largest dimension 2024-09-28 13:05:26 -06:00
Jaret Burkett
58537fc92b Added initial direct vision pixtral support 2024-09-28 10:47:51 -06:00
Jaret Burkett
86b5938cf3 Fixed the webp bug finally. 2024-09-25 13:56:00 -06:00
Jaret Burkett
6b4034122f REmove layers from direct vision resampler 2024-09-24 15:08:29 -06:00
Jaret Burkett
10817696fb Fixed issue where direct vision was not passing additional params from resampler when it is added 2024-09-24 10:34:11 -06:00
Jaret Burkett
037ce11740 Always return vision encoder in state dict 2024-09-24 07:43:17 -06:00
Jaret Burkett
04424fe2d6 Added config setting to set the timestep type 2024-09-24 06:53:59 -06:00
Jaret Burkett
40a8ff5731 Load local hugging face packages for assistant adapter 2024-09-23 10:37:12 -06:00
Jaret Burkett
2776221497 Added option to cache empty prompt or trigger and unload text encoders while training 2024-09-21 20:54:09 -06:00
Jaret Burkett
f85ad452c6 Added initial support for pixtral vision as a vision encoder. 2024-09-21 15:21:14 -06:00
Jaret Burkett
dd889086f4 Updates to the docker file for jupyterlab 2024-09-21 12:07:07 -06:00
apolinário
bc693488eb fix diffusers codebase (#183) 2024-09-21 11:50:29 -06:00
Jaret Burkett
d97c55cd96 Updated requirements to lock version of albucore, which had breaking changes. 2024-09-21 11:19:13 -06:00
Plat
79b4e04b80 Feat: Wandb logging (#95)
* wandb logging

* fix: start logging before train loop

* chore: add wandb dir to gitignore

* fix: wrap wandb functions

* fix: forget to send last samples

* chore: use valid type

* chore: use None when not type-checking

* chore: resolved complicated logic

* fix: follow log_every

---------

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

* small text changes

* no flash-attn? no problem!

* bye flash-attn!

* fixes for windows

---------

Co-authored-by: multimodalart <joaopaulo.passos+multimodal@gmail.com>
2024-08-30 06:29:51 -06:00
Jaret Burkett
60232def91 Made peleminary arch for flux ip adapter training 2024-08-28 08:55:39 -06:00
Jaret Burkett
3843e0d148 Added support for vision direct adapter for flux 2024-08-26 16:27:28 -06:00
liaoliaojun
e127c079da fix: print out the path where the image encode failed (#107) 2024-08-22 21:34:35 -06:00
martintomov
34db804c76 Modal cloud training support, fixed typo in toolkit/scheduler.py, Schnell training support for Colab, issue #92 , issue #114 (#115)
* issue #76, load_checkpoint_and_dispatch() 'force_hooks'

https://github.com/ostris/ai-toolkit/issues/76

* RunPod cloud config

https://github.com/ostris/ai-toolkit/issues/90

* change 2x A40 to 1x A40 and price per hour

referring to https://github.com/ostris/ai-toolkit/issues/90#issuecomment-2294894929

* include missed FLUX.1-schnell setup guide in last commit

* huggingface-cli login required auth

* #92 peft, #114 colab, schnell training in colab

* modal cloud - run_modal.py and .yaml configs

* run_modal.py mount path example

* modal_examples renamed to modal

* Training in Modal README.md setup guide

* rename run command in title for consistency
2024-08-22 21:25:44 -06:00
apolinário
4d35a29c97 Add push_to_hub to the trainer (#109)
* add push_to_hub

* fix indentation

* indent again

* model_config

* allow samples to not exist

* repo creation fix

* dont show empty [] if widget doesnt exist

* dont submit the config and optimizer

* Unsafe to have tokens saved in the yaml file

* make sure to catch only the latest samples

* change name to slug

* formatting

* formatting

---------

Co-authored-by: multimodalart <joaopaulo.passos+multimodal@gmail.com>
2024-08-22 21:18:56 -06:00
Jaret Burkett
b322d05fa3 Added tutorial link to readme 2024-08-22 16:25:32 -06:00
Jaret Burkett
8577849eeb Fixed wrong discord link. Woops. 2024-08-22 14:49:03 -06:00
Jaret Burkett
338c77d677 Fixed breaking change with diffusers. Allow flowmatch on normal stable diffusion models. 2024-08-22 14:36:22 -06:00
Jaret Burkett
e07a98a50c Bugfixes for full finetuning at bf16 2024-08-22 05:15:33 -06:00
Jaret Burkett
6a754b2710 Merge branch 'main' of github.com:ostris/ai-toolkit 2024-08-22 04:36:50 -06:00
Jaret Burkett
a939cf3730 WIP - adding support for flux DoRA and ip adapter training 2024-08-22 04:36:39 -06:00
Jaret Burkett
169dbd22ba Finaized bug reports 2024-08-18 16:21:48 -06:00
Jaret Burkett
6e7d721382 More issues testing 2024-08-18 16:20:08 -06:00
Jaret Burkett
dc6f36cd82 Testing github bug reporting stuff 2024-08-18 16:09:52 -06:00
martintomov
5603f9e004 issue #76, and RunPod cloud training setup #90 (#80)
* issue #76, load_checkpoint_and_dispatch() 'force_hooks'

https://github.com/ostris/ai-toolkit/issues/76

* RunPod cloud config

https://github.com/ostris/ai-toolkit/issues/90

* change 2x A40 to 1x A40 and price per hour

referring to https://github.com/ostris/ai-toolkit/issues/90#issuecomment-2294894929

* include missed FLUX.1-schnell setup guide in last commit

* huggingface-cli login required auth
2024-08-18 15:43:45 -06:00
Jaret Burkett
c45887192a Unload interum weights when doing multi lora fuse 2024-08-18 09:35:10 -06:00
Jaret Burkett
13a965a26c Fixed bad key naming on lora fuse I just pushed 2024-08-18 09:33:31 -06:00
Jaret Burkett
77ee7090e8 Update FAQ.md 2024-08-18 09:26:22 -06:00
Jaret Burkett
078396ceac Added a basic FAQ 2024-08-18 09:21:51 -06:00
Jaret Burkett
f944eeaa4d Fuse flux schnell assistant adapter in pieces when doing lowvram to drastically speed ip up from minutes to seconds. 2024-08-18 09:09:11 -06:00
Jaret Burkett
81899310f8 Added support for training on flux schnell. Added example config and instructions for training on flux schnell 2024-08-17 06:58:39 -06:00
Jaret Burkett
f9179540d2 Flush after sampling 2024-08-16 17:29:42 -06:00
Jaret Burkett
452e0e286d For lora assisted training, merge in before quantizing then sample with schnell at -1 weight. Almost doubles training speed with lora adapter. 2024-08-16 17:28:44 -06:00
Jaret Burkett
165510ace2 Dumb typo 2024-08-15 12:59:32 -06:00
Jaret Burkett
0355662e8e Added support for polarity guidance for flow matching models 2024-08-15 12:22:00 -06:00
Jaret Burkett
b99d36dfdb fixed issue with batch sizes larget than 1 2024-08-15 12:21:38 -06:00
Jaret Burkett
9001e5c933 Change flux latent spact if so it will not use old cache 2024-08-14 11:27:40 -06:00
Jaret Burkett
7fed4ea761 fixed huge flux training bug. Added ability to use an assistatn lora 2024-08-14 10:14:13 -06:00
Jaret Burkett
e07bf11727 Merge pull request #61 from fofr/patch-1
Fix image name in captions section of README
2024-08-14 08:01:51 -06:00
fofr
c728cc9a0b Update README.md 2024-08-14 15:00:02 +01:00
Jaret Burkett
00bd3d54a3 Actually use the save dtype from the config file. 2024-08-13 17:08:27 -06:00
Jaret Burkett
f7cf2f866f Make 100% sure lora alpha matches for flux 2024-08-13 14:24:03 -06:00
Jaret Burkett
465bc1e2f8 Update readme again 2024-08-13 13:37:22 -06:00
Jaret Burkett
0beca0d4a7 Updated readme 2024-08-13 13:35:20 -06:00
Jaret Burkett
418f5f7e8c Added new experimental time step weighing that should solve a lot of issues with distribution. Updated example. Removed a warning 2024-08-13 12:02:11 -06:00
Jaret Burkett
9ee1ef2a0a Added experimental modified sigma sqrt weight mapping for linear timestep scheduling for flowmatching 2024-08-12 17:03:09 -06:00
Jaret Burkett
599fafe01f Allow user to have the full flux checkpoint local 2024-08-12 09:57:16 -06:00
Jaret Burkett
af108bb964 Bug fix with dataloader. Added a flag to completly disable sampling 2024-08-12 09:19:40 -06:00
Jaret Burkett
89d61a3b8e Readme updates 2024-08-11 13:23:57 -06:00
Jaret Burkett
a6aa4b2c7d Added ability to set timesteps to linear for flowmatching schedule 2024-08-11 13:06:08 -06:00
Jaret Burkett
f8f0657b68 Added a colab notebook for training flux loras 2024-08-11 12:27:40 -06:00
Jaret Burkett
7f0ecdb377 Merge branch 'main' of github.com:ostris/ai-toolkit 2024-08-11 11:10:45 -06:00
Jaret Burkett
fbed8568fb Actually use the correct timestep sampling instead of calculating it and moving on lol. Tested a few with it and it seems to work better. 2024-08-11 11:10:37 -06:00
Jaret Burkett
6d31c6db73 Added a fix for windows dataloader 2024-08-11 10:48:24 -06:00
Jaret Burkett
6490a326e5 Fixed issue for vaes without a shift 2024-08-11 10:30:55 -06:00
Jaret Burkett
8d48ad4e85 fixed bug I added to demo config 2024-08-11 10:28:39 -06:00
Jaret Burkett
ec1ea7aa0e Added support for training on primary gpu with low_vram flag. Updated example script to remove creepy horse sample at that seed 2024-08-11 09:54:30 -06:00
Jaret Burkett
fa02e774b0 Added info about datset 2024-08-10 15:08:05 -06:00
Jaret Burkett
2308ef2868 Added flux training instructions 2024-08-10 14:10:02 -06:00
Jaret Burkett
b3e03295ad Reworked flux pred. Again 2024-08-08 13:06:34 -06:00
Jaret Burkett
e69a520616 Reworked timestep distribution on flowmatch sampler when training. 2024-08-08 06:01:45 -06:00
Jaret Burkett
acafe9984f Adjustments to loading of flux. Added a feedback to ema 2024-08-07 13:17:26 -06:00
Jaret Burkett
653fe60f16 Updates to flow matching algo 2024-08-07 15:04:17 +00:00
Jaret Burkett
c2424087d6 8 bit training working on flux 2024-08-06 11:53:27 -06:00
Jaret Burkett
272c8608c2 Make a CFG version of flux pipeline 2024-08-05 16:35:53 -06:00
Jaret Burkett
99f24cfb0c Added a conversion script to convert my loras to peft format for flux 2024-08-05 14:54:10 -06:00
Jaret Burkett
187663ab55 Use peft format for flux loras so they are compatible with diffusers. allow loading an assistant lora 2024-08-05 14:34:37 -06:00
Jaret Burkett
edb7e827ee Adjusted flow matching so target noise multiplier works properly with it. 2024-08-05 11:40:05 -06:00
Jaret Burkett
0ea27011d5 Bug fix 2024-08-04 11:07:19 -06:00
Jaret Burkett
f321de7bdb Setup to retrain guidance embedding for flux. Use defualt timestep distribution for flux 2024-08-04 10:37:23 -06:00
Jaret Burkett
88acc28d7f Prep for runpod docker 2024-08-03 12:41:06 -06:00
Jaret Burkett
de2da96a81 Updat4ed requirements 2024-08-03 09:50:18 -06:00
Jaret Burkett
9beea1c268 Flux training should work now... maybe 2024-08-03 09:17:34 -06:00
Jaret Burkett
369aa143bc Only train a few blocks on flux (for now) 2024-08-03 07:02:27 -06:00
Jaret Burkett
87ba867fdc Added flux training. Still a WIP. Wont train right without rectified flow working right 2024-08-02 15:00:30 -06:00
Jaret Burkett
03613c523f Bugfixes and cleanup 2024-08-01 11:45:12 -06:00
Jaret Burkett
47744373f2 Change img multiplier math 2024-07-30 11:33:41 -06:00
Jaret Burkett
443c996e7f Do a noisy unconsitional for vision direct 2024-07-29 15:42:26 -06:00
Jaret Burkett
8f0f467c20 Switch back to old ilora 2024-07-29 07:22:05 -06:00
Jaret Burkett
e81e19fd0f Added target_norm_std which is a game changer 2024-07-28 16:08:33 -06:00
Jaret Burkett
0bc4d555c7 A lot of pixart sigma training tweaks 2024-07-28 11:23:18 -06:00
Jaret Burkett
80aa2dbb80 New image generation img2img. various tweaks and fixes 2024-07-24 04:13:41 -06:00
Jaret Burkett
8d799031cf Remove reg as prior pred 2024-07-21 02:34:12 -06:00
Jaret Burkett
6e92922c14 Add a mergable linear to the mid of ilora 2024-07-20 21:17:53 -06:00
Jaret Burkett
c51235c486 Fixed misnamed var 2024-07-20 23:00:20 +00:00
Jaret Burkett
c2c4e8cf34 Added ability to target parts of lora for ilora 2024-07-20 22:45:52 +00:00
Jaret Burkett
4c249cf607 Added ilora2 2024-07-20 16:40:57 -06:00
Jaret Burkett
c2d5f712a3 Reworked ilora arch 2024-07-20 15:35:59 -06:00
Jaret Burkett
22d2f6e28f Fixed issue with grad scaling 2024-07-20 08:21:57 -06:00
Jaret Burkett
a2301cf28c Amall bug fixes 2024-07-18 10:39:55 -06:00
Jaret Burkett
11e426fdf1 Various features and fixes. Too much brain fog to do a proper description 2024-07-18 07:34:14 -06:00
Jaret Burkett
58dffd43a8 Added caching to image sizes so we dont do it every time. 2024-07-15 19:07:41 -06:00
Jaret Burkett
e4558dff4b Partial implementation for training auraflow. 2024-07-12 12:11:38 -06:00
Jaret Burkett
c062b7716c Varous bug fixes 2024-07-10 15:20:04 -06:00
Jaret Burkett
c008405480 Added after model load hook 2024-07-09 15:34:48 -06:00
Jaret Burkett
93e5df1d59 Merge branch 'main' of github.com:ostris/ai-toolkit 2024-07-07 07:56:56 -06:00
Jaret Burkett
045e4a6e15 Save entire pixart model again 2024-07-07 07:56:48 -06:00
Jaret Burkett
76f225a467 Fixed issue with TE adapter caption projection 2024-07-06 19:09:58 +00:00
Jaret Burkett
cab8a1c7b8 WIP to add the caption_proj weight to pixart sigma TE adapter 2024-07-06 13:00:21 -06:00
Jaret Burkett
acb06d6ff3 Bug fixes 2024-07-03 10:56:34 -06:00
Jaret Burkett
bb57623a35 Merge branch 'main' of github.com:ostris/ai-toolkit 2024-06-29 15:53:40 -06:00
Jaret Burkett
3072d20f17 Add ability to include conv_in and conv_out to full train when doing a lora 2024-06-29 14:54:50 -06:00
Jaret Burkett
f6b21f47bb Increased the number of heads for ip adapters. 2024-06-28 16:09:52 +00:00
Jaret Burkett
603ceca3ca added ema 2024-06-28 10:03:26 -06:00
Jaret Burkett
657fd09f25 Added more control over Sigma sizes 2024-06-26 08:57:53 -06:00
Jaret Burkett
8407c4deea Merge branch 'main' of github.com:ostris/ai-toolkit 2024-06-23 14:47:43 -06:00
Jaret Burkett
64f2b085b7 Minor fixes 2024-06-23 14:47:40 -06:00
Jaret Burkett
7165f2d25a Work to omprove pixart training 2024-06-23 20:46:48 +00:00
Jaret Burkett
5d47244c57 Added support for pixart sigma loras 2024-06-16 11:56:30 -06:00
Jaret Burkett
ada722c9e4 Fixed issue with heads not being added 2024-06-15 08:34:33 -06:00
Jaret Burkett
696f73c30d Removed variant 2024-06-14 17:09:47 -06:00
Jaret Burkett
e3410413b9 Rework head on ilora 2024-06-14 16:21:26 -06:00
Jaret Burkett
37cebd9458 WIP Ilora 2024-06-14 09:31:01 -06:00
Jaret Burkett
bd10d2d668 Some work on sd3 training. Not working 2024-06-13 12:19:16 -06:00
Jaret Burkett
cb5d28cba9 Added working ilora trainer 2024-06-12 09:33:45 -06:00
Jaret Burkett
3f3636b788 Bug fixes and little improvements here and there. 2024-06-08 06:24:20 -06:00
Jaret Burkett
833c833f28 WIP on SAFE encoder. Work on fp16 training improvements. Various other tweaks and improvements 2024-05-27 10:50:24 -06:00
Jaret Burkett
68b7e159bc Bug Fixes 2024-05-17 08:41:20 -06:00
Jaret Burkett
5a45c709cd Work on ipadapters and custom adapters 2024-05-13 06:37:54 -06:00
Jaret Burkett
10e1ecf1e8 Added single value adapter training 2024-04-28 06:04:47 -06:00
Jaret Burkett
b96913d73c Improvements to dataloader 2024-04-27 09:28:28 -06:00
Jaret Burkett
5da3613e0b Bug fixes and minor features 2024-04-25 06:14:31 -06:00
Jaret Burkett
5a70b7f38d Added pixart sigma support, but it wont work until i address breaking changes with lora code in diffusers so it can be upgraded. 2024-04-20 10:46:56 -06:00
Jaret Burkett
377b81ee3e Adjustments to guidance 2024-04-19 15:00:35 -06:00
Jaret Burkett
2d0a1be59d Bug fixes 2024-04-16 03:48:13 -06:00
Jaret Burkett
7284aab7c0 Added specialized scaler training to ip adapters 2024-04-05 08:17:09 -06:00
Jaret Burkett
427847ac4c Small tweaks and fixes for specialized ip adapter training 2024-03-26 11:35:26 -06:00
Jaret Burkett
9c1cc9641e Added keep tokens to keep so many tokens in a prompt when dropping 2024-03-18 13:18:25 -06:00
Jaret Burkett
89f4bcad2e Lock diffusers to 0.26.3 until I can figure out why future versions break LoRA code 2024-03-18 10:17:55 -06:00
Jaret Burkett
016687bda1 Adapter work. Bug fixes. Auto adjust LR when resuming optimizer. 2024-03-17 10:21:47 -06:00
Jaret Burkett
72de68d8aa WIP on clip vision encoder 2024-03-13 07:24:08 -06:00
Jaret Burkett
d87b49882c Work on embedding adapters 2024-03-11 15:18:42 -06:00
Jaret Burkett
f415bac7b5 Merge branch 'main' of github.com:ostris/ai-toolkit 2024-03-06 09:32:38 -07:00
Jaret Burkett
f1cb87fe9e fixed bug the kept learning rates the same 2024-03-06 09:23:32 -07:00
Jaret Burkett
8f9cd823d1 Create LICENSE 2024-03-06 07:54:55 -07:00
Jaret Burkett
b01e8d889a Added stochastic rounding to adafactor. ILora adjustments 2024-03-05 07:07:09 -07:00
Jaret Burkett
1325613583 rework ilora 2024-02-29 07:55:52 -07:00
Jaret Burkett
337945de9a Added this not that guidance. Added ability to replace prompts. 2024-02-28 20:10:14 -07:00
Jaret Burkett
561914d8e6 Removed old code for fixing multistep sampler that is no longer needed 2024-02-25 11:53:35 -07:00
Jaret Burkett
b0a0f28191 Bug fixes 2024-02-25 08:28:29 -07:00
Jaret Burkett
f965a1299f Fixed Dora implementation. Still highly experimental 2024-02-24 10:26:01 -07:00
Jaret Burkett
1bd94f0f01 Added early DoRA support, but will change shortly. Dont use right now. 2024-02-23 05:55:41 -07:00
Jaret Burkett
9ffa8c3711 Fixed issue when there is no adapter 2024-02-22 02:59:59 -07:00
Jaret Burkett
b68c3ef734 Added te aug adapter 2024-02-21 21:30:26 -07:00
Jaret Burkett
49c41e6a5f Bug fixes. allow for random negative prompts 2024-02-21 04:51:52 -07:00
Jaret Burkett
2478554c95 Bug fixes. Added IP adapter training for Pixart 2024-02-17 10:06:57 -07:00
Jaret Burkett
93b52932c1 Added training for pixart-a 2024-02-13 16:00:04 -07:00
Jaret Burkett
4ec4025cbb Added adapter modules for text encoders and direct vision 2024-02-12 08:46:18 -07:00
Jaret Burkett
e074058faa Work on additional image embedding methods. Finalized zipper resampler. It works amazing 2024-02-10 09:00:05 -07:00
Jaret Burkett
a8481c1670 randomly adjust scale of unconditional noise on ip adapters if training with cfg 2024-02-06 03:44:54 -07:00
Jaret Burkett
e18e0cb5f8 Added comparitive loss when training clip encoder. Allow selecting clip layer. on ip adapter. Improvements to prior prediction 2024-02-05 07:40:03 -07:00
Jaret Burkett
177c7130ec improved correction of pred norm by targeting the prior 2024-02-01 06:31:04 -07:00
Jaret Burkett
1ae1017748 Bug fixes. added ability to use l1 loss. varous other tests and improvements 2024-01-31 06:30:54 -07:00
Jaret Burkett
92b9c71d44 Many bug fixes. Ip adapter bug fixes. Added noise to unconditional, it works better. added an ilora adapter for 1 shotting LoRAs 2024-01-28 08:20:03 -07:00
Jaret Burkett
f17ad8d794 various bug fixes. Created an contextual alpha mask module to calculate alpha mask 2024-01-18 16:34:27 -07:00
Jaret Burkett
86c70a2a1f Added an experimental clip fusion model that is showing promise for embedding concepts 2024-01-17 13:13:04 -07:00
Jaret Burkett
655533d4c7 More work on custom adapter 2024-01-16 17:41:26 -07:00
Jaret Burkett
eebd3c8212 Initial training script for photomaker training. Needs a little more work. 2024-01-15 18:46:26 -07:00
Jaret Burkett
5276975fb0 Added additional config options for custom plugins I needed 2024-01-15 08:31:09 -07:00
Jaret Burkett
e190fbaeb8 Prepwork for ilora 2024-01-12 06:41:15 -07:00
Jaret Burkett
290393f7ae Imporvements to ip weight adaptation. Bug fixes. Added masking to direct guidance loss. Allow importing a file for random triggers. Handle bas meta images with improper sizing. 2024-01-11 12:22:16 -07:00
Jaret Burkett
b2a54c8f36 Added siglip support 2024-01-09 20:52:21 -07:00
Jaret Burkett
b767d29b3c Adjustments to the clip preprocessor. Allow merging in new weights for ip adapters so you can change the arcitecture while maintaining as much data as possible 2024-01-06 11:56:53 -07:00
Jaret Burkett
645b27f97a Bug fixes with ip adapter training. Made a clip pre processor that can be trained with ip adapter to help augment the clip input to squeeze in more detail from a larget input. moved clip processing to the dataloader for speed. 2024-01-04 12:59:38 -07:00
Jaret Burkett
65c08b09c3 Added ability to do cfg during training. Various bug fixes 2024-01-02 11:29:57 -07:00
Jaret Burkett
afc231efc1 Added reference adapters, many bug fixes, more ip adapter work and customizability 2024-01-01 17:15:53 -07:00
Jaret Burkett
bafacf3b65 Initial commit 2023-12-29 13:07:35 -07:00
Jaret Burkett
0892dec4a5 Fixed some new bugs i added. woops 2023-12-28 14:03:42 -07:00
Jaret Burkett
eeee4a1620 Created a size agnostic feature encoder (SAFE) model to be trained in replace of CLIP for ip adapters. It is mostly conv layers so will hopefully be able to handle facial features better than clip can. Also bug fixes 2023-12-28 12:20:27 -07:00
Jaret Burkett
d11ed7f66c Big fixes and added method to standardize values in both latent and pixel space before feeding into the network. Target values were determined over huge generated regularization sets. 2023-12-26 06:19:48 -07:00
Jaret Burkett
27ad79053e Added SDXL support for clip vision embedder trainer 2023-12-24 14:31:29 -07:00
Jaret Burkett
05ae95ca89 Added a clip vision adapter trainer. Only works for sd15 for now 2023-12-24 13:26:04 -07:00
Jaret Burkett
0f8daa5612 Bug fixes, work on maing IP adapters more customizable. 2023-12-24 08:32:39 -07:00
Jaret Burkett
7703e3a15e Fixes for sdxl ip adapter training. Bug fixes 2023-12-21 11:15:58 -07:00
Jaret Burkett
0f597f453e Switched ip adapter dataloader to clip_image paths so the control paths can be used for training assistant adapters while training ip adapters 2023-12-20 10:32:24 -07:00
Jaret Burkett
dfb64b5957 Allow ip adapters to be much more variable in their creation 2023-12-20 06:18:33 -07:00
Jaret Burkett
82098e5d6e Added more functionality for ip adapters 2023-12-19 09:54:56 -07:00
Jaret Burkett
b653906715 Fixed ip adapter training. Works now 2023-12-17 08:22:59 -07:00
Jaret Burkett
13d32423f6 Added a polarity balancer to guidance 2023-12-15 15:19:14 -07:00
Jaret Burkett
39870411d8 More guidance work. Improved LoRA module resolver for unet. Added vega mappings and LoRA training for it. Various other bigfixes and changes 2023-12-15 06:02:10 -07:00
Jaret Burkett
e5177833b2 Targeted guidance work 2023-12-09 19:06:18 -07:00
Jaret Burkett
eaa0fb6253 Tons of bug fixes and improvements to special training. Fixed slider training. 2023-12-09 16:38:10 -07:00
Jaret Burkett
eaec2f5a52 Added guidance mentiods. WIP 2023-12-08 08:39:21 -07:00
Jaret Burkett
92cb5ae096 Reworked targeted guidance algo 2023-12-01 06:30:52 -07:00
Jaret Burkett
bd2bce9b92 Switched to trailing timestep spacing to make timesteps for consistant across schedulers. Honed in on targeted guidance. It is finally perfect. (I think) 2023-11-29 14:32:48 -07:00
Jaret Burkett
537af79b0d Merge branch 'main' of github.com:ostris/ai-toolkit 2023-11-29 10:13:55 -07:00
Jaret Burkett
7624241032 More fixes for noise schedules and fixed targeted guidance inverted masked prior 2023-11-29 10:13:31 -07:00
Jaret Burkett
0d5943af91 Update readme for torch requirements 2023-11-28 12:52:55 -07:00
Jaret Burkett
be815f9c47 updated requirements 2023-11-28 10:43:17 -07:00
Jaret Burkett
3443d6aafa Updated Readme 2023-11-28 10:41:16 -07:00
Jaret Burkett
bef10a639c Merge remote-tracking branch 'origin/development'
# Conflicts:
#	toolkit/stable_diffusion_model.py
2023-11-28 10:40:05 -07:00
Jaret Burkett
3eb3535683 Merge pull request #12 from bendeguzvaradi/main
Bug/Safety checker to None
2023-09-14 15:31:09 -06:00
bendeguzvaradi
3d387103cd safety checker to None 2023-09-11 17:05:55 +02:00
237 changed files with 45096 additions and 4617 deletions

2
.github/FUNDING.yml vendored Normal file
View File

@@ -0,0 +1,2 @@
github: [ostris]
patreon: ostris

19
.github/ISSUE_TEMPLATE/bug_report.md vendored Normal file
View File

@@ -0,0 +1,19 @@
---
name: Bug Report
about: For bugs only. Not for feature requests or questions.
title: ''
labels: ''
assignees: ''
---
## This is for bugs only
Did you already ask [in the discord](https://discord.gg/VXmU2f5WEU)?
Yes/No
You verified that this is a bug and not a feature request or question by asking [in the discord](https://discord.gg/VXmU2f5WEU)?
Yes/No
## Describe the bug

5
.github/ISSUE_TEMPLATE/config.yml vendored Normal file
View File

@@ -0,0 +1,5 @@
blank_issues_enabled: false
contact_links:
- name: Ask in the Discord BEFORE opening an issue
url: https://discord.gg/VXmU2f5WEU
about: Please ask in the discord before opening a github issue.

9
.gitignore vendored
View File

@@ -161,6 +161,7 @@ cython_debug/
/env.sh
/models
/datasets
/custom/*
!/custom/.gitkeep
/.tmp
@@ -172,4 +173,10 @@ cython_debug/
/output/*
!/output/.gitkeep
/extensions/*
!/extensions/example
!/extensions/example
/temp
/wandb
.vscode/settings.json
.DS_Store
._.DS_Store
aitk_db.db

4
.gitmodules vendored
View File

@@ -1,12 +1,16 @@
[submodule "repositories/sd-scripts"]
path = repositories/sd-scripts
url = https://github.com/kohya-ss/sd-scripts.git
commit = b78c0e2a69e52ce6c79abc6c8c82d1a9cabcf05c
[submodule "repositories/leco"]
path = repositories/leco
url = https://github.com/p1atdev/LECO
commit = 9294adf40218e917df4516737afb13f069a6789d
[submodule "repositories/batch_annotator"]
path = repositories/batch_annotator
url = https://github.com/ostris/batch-annotator
commit = 420e142f6ad3cc14b3ea0500affc2c6c7e7544bf
[submodule "repositories/ipadapter"]
path = repositories/ipadapter
url = https://github.com/tencent-ailab/IP-Adapter.git
commit = 5a18b1f3660acaf8bee8250692d6fb3548a19b14

28
.vscode/launch.json vendored Normal file
View File

@@ -0,0 +1,28 @@
{
"version": "0.2.0",
"configurations": [
{
"name": "Run current config",
"type": "python",
"request": "launch",
"program": "${workspaceFolder}/run.py",
"args": [
"${file}"
],
"env": {
"CUDA_LAUNCH_BLOCKING": "1",
"DEBUG_TOOLKIT": "1"
},
"console": "integratedTerminal",
"justMyCode": false
},
{
"name": "Python: Debug Current File",
"type": "python",
"request": "launch",
"program": "${file}",
"console": "integratedTerminal",
"justMyCode": false
},
]
}

10
FAQ.md Normal file
View File

@@ -0,0 +1,10 @@
# FAQ
WIP. Will continue to add things as they are needed.
## FLUX.1 Training
#### How much VRAM is required to train a lora on FLUX.1?
24GB minimum is required.

21
LICENSE Normal file
View File

@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 Ostris, LLC
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.

444
README.md

File diff suppressed because one or more lines are too long

40
assets/glif.svg Normal file
View File

@@ -0,0 +1,40 @@
<svg width="148" height="66" viewBox="0 0 148 66" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect opacity="0.3" width="148" height="66" rx="33" fill="#030F2F"/>
<g filter="url(#filter0_d_10631_12135)">
<path d="M48.8305 21.013H43.5433V53.7839H48.8305V21.013Z" fill="white"/>
<path d="M57.0987 32.7729H51.8115V53.7835H57.0987V32.7729Z" fill="white"/>
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
<path d="M58.4255 24.7602L54.4794 20.7832L50.5333 24.7602L54.4794 28.7372L58.4255 24.7602Z" fill="white"/>
<path d="M40.5557 21.0118H35.2685V24.185C33.9942 23.5456 32.5588 23.1832 31.0387 23.1832C25.7911 23.1806 21.5217 27.4834 21.5217 32.7721C21.5217 35.082 22.336 37.2054 23.6921 38.8626H21.5217V44.1912H26.8089V41.3618C28.0832 42.0012 29.5186 42.3635 31.0387 42.3635C36.2863 42.3635 40.5557 38.0607 40.5557 32.7721C40.5557 30.2996 39.6225 28.0429 38.0918 26.3404H40.5557V21.0118ZM31.0387 37.0349C28.707 37.0349 26.8089 35.122 26.8089 32.7721C26.8089 30.4221 28.707 28.5092 31.0387 28.5092C33.3704 28.5092 35.2685 30.4221 35.2685 32.7721C35.2685 35.122 33.3704 37.0349 31.0387 37.0349Z" fill="white"/>
<path d="M31.0381 44.1912H26.8083V49.5198H31.0381C33.3697 49.5198 35.2678 51.4327 35.2678 53.7826H40.555C40.555 48.494 36.2856 44.1912 31.0381 44.1912Z" fill="white"/>
<path d="M69.6041 26.3416C71.1295 26.3416 72.6073 27.1702 73.3871 28.6968L73.596 28.4864L77.2098 24.8443C75.4253 22.4544 72.6126 21.013 69.6041 21.013C64.4438 21.013 60.2352 25.172 60.0951 30.338L60.0872 53.7839H65.3744V30.6018C65.3744 28.2519 67.2725 26.3389 69.6041 26.3389V26.3416Z" fill="white"/>
<path d="M73.5495 36.4837L69.6034 32.5067L65.6573 36.4837L69.6034 40.4607L73.5495 36.4837Z" fill="white"/>
<path d="M120.022 53.8259H117.218V32.6354H120.022V35.2219C121.02 33.321 122.702 32.2615 125.102 32.2615C129.371 32.2615 131.397 35.9698 131.397 40.5819C131.397 45.1939 129.371 48.9023 125.102 48.9023C122.702 48.9023 121.02 47.8427 120.022 45.9418V53.8259ZM120.022 39.2419V41.9219C120.022 44.6642 121.581 46.6586 124.385 46.6586C126.722 46.6586 128.436 45.1939 128.436 43.0437V38.12C128.436 35.9698 126.722 34.5052 124.385 34.5052C121.581 34.5052 120.022 36.4996 120.022 39.2419Z" fill="white"/>
<path d="M103.267 53.8259H100.463V32.6354H103.267V35.2219C104.265 33.321 105.947 32.2615 108.347 32.2615C112.616 32.2615 114.642 35.9698 114.642 40.5819C114.642 45.1939 112.616 48.9023 108.347 48.9023C105.947 48.9023 104.265 47.8427 103.267 45.9418V53.8259ZM103.267 39.2419V41.9219C103.267 44.6642 104.826 46.6586 107.63 46.6586C109.967 46.6586 111.681 45.1939 111.681 43.0437V38.12C111.681 35.9698 109.967 34.5052 107.63 34.5052C104.826 34.5052 103.267 36.4996 103.267 39.2419Z" fill="white"/>
<path d="M87.7844 48.9023C86.2262 48.9023 84.8862 48.4037 83.9825 47.4688C83.1723 46.6274 82.6737 45.4121 82.6737 44.1656C82.6737 41.6726 84.263 39.6782 87.4104 39.3977L92.5834 38.9303V37.6526C92.5834 35.1907 91.2434 34.5052 89.2802 34.5052C87.3169 34.5052 86.0081 35.3466 86.0081 37.3721H83.2035C83.2035 34.3805 85.9146 32.2615 89.3113 32.2615C92.7392 32.2615 95.3257 33.9442 95.3257 37.6838V46.1288H97.694V48.5283H92.9573L92.895 45.599H92.7704C91.8978 47.8427 90.1527 48.9023 87.7844 48.9023ZM88.5011 46.6586C91.0253 46.6586 92.5834 44.6642 92.5834 41.7972V41.0805L86.943 41.6102C85.79 41.7037 85.6341 42.2335 85.6341 43.1372V44.5395C85.6341 46.0041 86.7248 46.6586 88.5011 46.6586Z" fill="white"/>
<path d="M80.5279 45.1002V48.5281H77.1V45.1002H80.5279Z" fill="white"/>
</g>
<path d="M102.683 12.3401L100.874 9.09521H101.922L102.976 11.0436C103.015 11.1169 103.05 11.1852 103.079 11.2487C103.108 11.3122 103.138 11.3757 103.167 11.4392C103.191 11.3952 103.211 11.3537 103.225 11.3147C103.24 11.2756 103.257 11.2365 103.277 11.1975C103.301 11.1535 103.328 11.1022 103.357 11.0436L104.405 9.09521H105.423L103.621 12.3401V14.4497H102.683V12.3401Z" fill="white"/>
<path d="M97.1749 9.09521V14.4497H96.2373V9.09521H97.1749ZM98.3615 12.1717H96.8892V11.3806H98.3102C98.5691 11.3806 98.7668 11.3171 98.9036 11.1901C99.0403 11.0583 99.1087 10.8727 99.1087 10.6334C99.1087 10.4039 99.0378 10.2281 98.8962 10.106C98.7546 9.98397 98.5495 9.92292 98.2809 9.92292H96.8599V9.09521H98.3615C98.8938 9.09521 99.3113 9.22462 99.6141 9.48343C99.9168 9.74224 100.068 10.0963 100.068 10.5455C100.068 10.8678 99.9901 11.1389 99.8338 11.3586C99.6776 11.5735 99.4456 11.7297 99.1379 11.8274V11.7248C99.47 11.803 99.7215 11.9495 99.8924 12.1643C100.063 12.3792 100.149 12.6575 100.149 12.9994C100.149 13.3021 100.08 13.5634 99.9437 13.7831C99.807 13.998 99.6067 14.164 99.343 14.2812C99.0842 14.3935 98.7717 14.4497 98.4055 14.4497H96.8599V13.622H98.3615C98.6301 13.622 98.8352 13.5585 98.9768 13.4315C99.1184 13.3046 99.1892 13.1215 99.1892 12.8822C99.1892 12.6575 99.116 12.4842 98.9695 12.3621C98.8279 12.2351 98.6252 12.1717 98.3615 12.1717Z" fill="white"/>
<path d="M89.5954 14.4497H87.6689V9.09521H89.5441C90.0715 9.09521 90.5354 9.20997 90.9358 9.43948C91.3363 9.66411 91.6488 9.97908 91.8734 10.3844C92.1029 10.7848 92.2177 11.2512 92.2177 11.7834C92.2177 12.3059 92.1054 12.7699 91.8807 13.1752C91.661 13.5756 91.3534 13.8881 90.9578 14.1128C90.5671 14.3374 90.113 14.4497 89.5954 14.4497ZM88.6065 9.52738V14.0249L88.1597 13.5854H89.5075C89.864 13.5854 90.1716 13.5121 90.4304 13.3656C90.6892 13.2191 90.887 13.0116 91.0237 12.743C91.1605 12.4744 91.2288 12.1546 91.2288 11.7834C91.2288 11.4025 91.158 11.0778 91.0164 10.8092C90.8748 10.5358 90.6721 10.3258 90.4084 10.1793C90.1447 10.0328 89.8273 9.95955 89.4562 9.95955H88.1597L88.6065 9.52738Z" fill="white"/>
<path d="M86.0735 14.4497H82.748V9.09521H86.0735V9.95955H83.356L83.6856 9.65923V11.3366H85.8245V12.1643H83.6856V13.8857L83.356 13.5854H86.0735V14.4497Z" fill="white"/>
<path d="M78.1926 14.4497H77.255V9.09521H79.2986C79.9042 9.09521 80.3754 9.24171 80.7123 9.53471C81.0542 9.8277 81.2251 10.2379 81.2251 10.7653C81.2251 11.1218 81.1421 11.427 80.976 11.6809C80.8149 11.9299 80.5756 12.1204 80.2582 12.2522L81.2764 14.4497H80.2509L79.3426 12.45H78.1926V14.4497ZM78.1926 9.93025V11.6223H79.2986C79.5965 11.6223 79.8285 11.5466 79.9945 11.3952C80.1605 11.2438 80.2436 11.0339 80.2436 10.7653C80.2436 10.4967 80.1605 10.2916 79.9945 10.15C79.8285 10.0035 79.5965 9.93025 79.2986 9.93025H78.1926Z" fill="white"/>
<path d="M75.789 11.7688C75.789 12.3108 75.6792 12.7918 75.4594 13.2118C75.2397 13.6269 74.9345 13.9516 74.5438 14.186C74.1531 14.4204 73.7014 14.5376 73.1887 14.5376C72.6808 14.5376 72.2316 14.4204 71.8409 14.186C71.4503 13.9516 71.1451 13.6269 70.9253 13.2118C70.7105 12.7967 70.603 12.3182 70.603 11.7761C70.603 11.2292 70.7129 10.7482 70.9326 10.3331C71.1524 9.91317 71.4552 9.58599 71.8409 9.35159C72.2316 9.1172 72.6833 9 73.196 9C73.7088 9 74.158 9.1172 74.5438 9.35159C74.9345 9.58599 75.2397 9.91073 75.4594 10.3258C75.6792 10.7409 75.789 11.2219 75.789 11.7688ZM74.8075 11.7688C74.8075 11.3879 74.7416 11.0583 74.6097 10.7799C74.4779 10.5016 74.2923 10.2867 74.053 10.1354C73.8138 9.97909 73.5281 9.90096 73.196 9.90096C72.8689 9.90096 72.5832 9.97909 72.339 10.1354C72.0997 10.2867 71.9142 10.5016 71.7823 10.7799C71.6505 11.0583 71.5846 11.3879 71.5846 11.7688C71.5846 12.1497 71.6505 12.4818 71.7823 12.765C71.9142 13.0433 72.0997 13.2582 72.339 13.4096C72.5832 13.561 72.8689 13.6366 73.196 13.6366C73.5281 13.6366 73.8138 13.561 74.053 13.4096C74.2923 13.2533 74.4779 13.036 74.6097 12.7577C74.7416 12.4744 74.8075 12.1448 74.8075 11.7688Z" fill="white"/>
<path d="M65.6821 10.5895C65.6821 10.277 65.7627 10.0011 65.9239 9.76179C66.085 9.52251 66.3072 9.33694 66.5904 9.2051C66.8785 9.06837 67.2106 9 67.5866 9C67.948 9 68.2605 9.06348 68.5242 9.19045C68.7928 9.31741 69.0003 9.49809 69.1468 9.73249C69.2982 9.96688 69.3788 10.2452 69.3885 10.5675H68.4509C68.4412 10.338 68.3582 10.1598 68.2019 10.0328C68.0456 9.90096 67.8356 9.83504 67.572 9.83504C67.2838 9.83504 67.0519 9.90096 66.8761 10.0328C66.7052 10.1598 66.6197 10.3356 66.6197 10.5602C66.6197 10.7506 66.671 10.902 66.7735 11.0143C66.881 11.1218 67.047 11.2023 67.2716 11.2561L68.114 11.4465C68.573 11.5442 68.9148 11.7126 69.1395 11.9519C69.3641 12.1863 69.4764 12.5037 69.4764 12.9042C69.4764 13.2313 69.3958 13.5194 69.2347 13.7685C69.0736 14.0175 68.844 14.2104 68.5462 14.3472C68.2532 14.479 67.9089 14.5449 67.5134 14.5449C67.1373 14.5449 66.8077 14.4814 66.5245 14.3545C66.2413 14.2226 66.0191 14.0395 65.8579 13.8051C65.7017 13.5707 65.6187 13.2948 65.6089 12.9774H66.5465C66.5514 13.202 66.6393 13.3803 66.8102 13.5121C66.986 13.6391 67.2228 13.7026 67.5207 13.7026C67.8332 13.7026 68.0798 13.6391 68.2605 13.5121C68.4461 13.3803 68.5388 13.2069 68.5388 12.9921C68.5388 12.8065 68.49 12.66 68.3923 12.5526C68.2947 12.4402 68.136 12.3621 67.9162 12.3182L67.0665 12.1277C66.6124 12.0301 66.2681 11.8543 66.0337 11.6003C65.7993 11.3415 65.6821 11.0046 65.6821 10.5895Z" fill="white"/>
<path d="M60.8331 14.4497H59.9102V9.09521H60.8404L63.6239 13.307H63.3528V9.09521H64.2758V14.4497H63.3528L60.5621 10.2452H60.8331V14.4497Z" fill="white"/>
<path d="M58.4443 11.7688C58.4443 12.3108 58.3344 12.7918 58.1147 13.2118C57.8949 13.6269 57.5897 13.9516 57.1991 14.186C56.8084 14.4204 56.3567 14.5376 55.844 14.5376C55.3361 14.5376 54.8869 14.4204 54.4962 14.186C54.1055 13.9516 53.8003 13.6269 53.5806 13.2118C53.3657 12.7967 53.2583 12.3182 53.2583 11.7761C53.2583 11.2292 53.3682 10.7482 53.5879 10.3331C53.8077 9.91317 54.1104 9.58599 54.4962 9.35159C54.8869 9.1172 55.3386 9 55.8513 9C56.364 9 56.8133 9.1172 57.1991 9.35159C57.5897 9.58599 57.8949 9.91073 58.1147 10.3258C58.3344 10.7409 58.4443 11.2219 58.4443 11.7688ZM57.4628 11.7688C57.4628 11.3879 57.3969 11.0583 57.265 10.7799C57.1332 10.5016 56.9476 10.2867 56.7083 10.1354C56.469 9.97909 56.1834 9.90096 55.8513 9.90096C55.5241 9.90096 55.2385 9.97909 54.9943 10.1354C54.755 10.2867 54.5695 10.5016 54.4376 10.7799C54.3058 11.0583 54.2398 11.3879 54.2398 11.7688C54.2398 12.1497 54.3058 12.4818 54.4376 12.765C54.5695 13.0433 54.755 13.2582 54.9943 13.4096C55.2385 13.561 55.5241 13.6366 55.8513 13.6366C56.1834 13.6366 56.469 13.561 56.7083 13.4096C56.9476 13.2533 57.1332 13.036 57.265 12.7577C57.3969 12.4744 57.4628 12.1448 57.4628 11.7688Z" fill="white"/>
<path d="M49.2393 9.09521V14.4497H48.3018V9.09521H49.2393ZM50.4186 12.6038H49.0123V11.7688H50.2209C50.5432 11.7688 50.7873 11.6882 50.9534 11.5271C51.1243 11.361 51.2097 11.1315 51.2097 10.8385C51.2097 10.5455 51.1243 10.3209 50.9534 10.1646C50.7873 10.0084 50.5481 9.93025 50.2355 9.93025H48.9244V9.09521H50.4186C50.78 9.09521 51.0925 9.16846 51.3562 9.31496C51.6199 9.46146 51.825 9.66655 51.9715 9.93025C52.118 10.1891 52.1913 10.4943 52.1913 10.8459C52.1913 11.1877 52.118 11.4929 51.9715 11.7615C51.825 12.0252 51.6199 12.2327 51.3562 12.3841C51.0925 12.5306 50.78 12.6038 50.4186 12.6038Z" fill="white"/>
<path d="M43.0732 10.5895C43.0732 10.277 43.1538 10.0011 43.315 9.76179C43.4761 9.52251 43.6983 9.33694 43.9815 9.2051C44.2696 9.06837 44.6017 9 44.9777 9C45.3391 9 45.6516 9.06348 45.9153 9.19045C46.1839 9.31741 46.3914 9.49809 46.5379 9.73249C46.6893 9.96688 46.7699 10.2452 46.7796 10.5675H45.842C45.8323 10.338 45.7493 10.1598 45.593 10.0328C45.4367 9.90096 45.2268 9.83504 44.9631 9.83504C44.675 9.83504 44.443 9.90096 44.2672 10.0328C44.0963 10.1598 44.0108 10.3356 44.0108 10.5602C44.0108 10.7506 44.0621 10.902 44.1647 11.0143C44.2721 11.1218 44.4381 11.2023 44.6627 11.2561L45.5051 11.4465C45.9641 11.5442 46.306 11.7126 46.5306 11.9519C46.7552 12.1863 46.8675 12.5037 46.8675 12.9042C46.8675 13.2313 46.787 13.5194 46.6258 13.7685C46.4647 14.0175 46.2352 14.2104 45.9373 14.3472C45.6443 14.479 45.3 14.5449 44.9045 14.5449C44.5285 14.5449 44.1988 14.4814 43.9156 14.3545C43.6324 14.2226 43.4102 14.0395 43.249 13.8051C43.0928 13.5707 43.0098 13.2948 43 12.9774H43.9376C43.9425 13.202 44.0304 13.3803 44.2013 13.5121C44.3771 13.6391 44.6139 13.7026 44.9118 13.7026C45.2243 13.7026 45.4709 13.6391 45.6516 13.5121C45.8372 13.3803 45.9299 13.2069 45.9299 12.9921C45.9299 12.8065 45.8811 12.66 45.7835 12.5526C45.6858 12.4402 45.5271 12.3621 45.3073 12.3182L44.4576 12.1277C44.0035 12.0301 43.6592 11.8543 43.4248 11.6003C43.1904 11.3415 43.0732 11.0046 43.0732 10.5895Z" fill="white"/>
<defs>
<filter id="filter0_d_10631_12135" x="20.0869" y="19.3484" width="112.745" height="35.9123" filterUnits="userSpaceOnUse" color-interpolation-filters="sRGB">
<feFlood flood-opacity="0" result="BackgroundImageFix"/>
<feColorMatrix in="SourceAlpha" type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 127 0" result="hardAlpha"/>
<feOffset/>
<feGaussianBlur stdDeviation="0.717391"/>
<feComposite in2="hardAlpha" operator="out"/>
<feColorMatrix type="matrix" values="0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0.25 0"/>
<feBlend mode="normal" in2="BackgroundImageFix" result="effect1_dropShadow_10631_12135"/>
<feBlend mode="normal" in="SourceGraphic" in2="effect1_dropShadow_10631_12135" result="shape"/>
</filter>
</defs>
</svg>

After

Width:  |  Height:  |  Size: 13 KiB

BIN
assets/lora_ease_ui.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 340 KiB

29
build_and_push_docker Normal file
View File

@@ -0,0 +1,29 @@
#!/usr/bin/env bash
# Extract version from version.py
if [ -f "version.py" ]; then
VERSION=$(python3 -c "from version import VERSION; print(VERSION)")
echo "Building version: $VERSION"
else
echo "Error: version.py not found. Please create a version.py file with VERSION defined."
exit 1
fi
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
echo "Building version: $VERSION and latest"
# wait 2 seconds
sleep 2
# Build the image with cache busting
docker build --build-arg CACHEBUST=$(date +%s) -t aitoolkit:$VERSION -f docker/Dockerfile .
# Tag with version and latest
docker tag aitoolkit:$VERSION ostris/aitoolkit:$VERSION
docker tag aitoolkit:$VERSION ostris/aitoolkit:latest
# Push both tags
echo "Pushing images to Docker Hub..."
docker push ostris/aitoolkit:$VERSION
docker push ostris/aitoolkit:latest
echo "Successfully built and pushed ostris/aitoolkit:$VERSION and ostris/aitoolkit:latest"

View File

@@ -0,0 +1,96 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flux_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
- folder_path: "/root/ai-toolkit/your-dataset"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation_steps: 1
train_unet: true
train_text_encoder: false # probably won't work with flux
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new vell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
# if you get an error, or get stuck while downloading,
# check https://github.com/ostris/ai-toolkit/issues/84, download the model locally and
# place it like "/root/ai-toolkit/FLUX.1-dev"
name_or_path: "black-forest-labs/FLUX.1-dev"
is_flux: true
quantize: true # run 8bit mixed precision
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flux
seed: 42
walk_seed: true
guidance_scale: 4
sample_steps: 20
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,98 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flux_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "/root/ai-toolkit/modal_output" # must match MOUNT_DIR from run_modal.py
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
# your dataset must be placed in /ai-toolkit and /root is for modal to find the dir:
- folder_path: "/root/ai-toolkit/your-dataset"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation_steps: 1
train_unet: true
train_text_encoder: false # probably won't work with flux
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new vell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
# if you get an error, or get stuck while downloading,
# check https://github.com/ostris/ai-toolkit/issues/84, download the models locally and
# place them like "/root/ai-toolkit/FLUX.1-schnell" and "/root/ai-toolkit/FLUX.1-schnell-training-adapter"
name_or_path: "black-forest-labs/FLUX.1-schnell"
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
is_flux: true
quantize: true # run 8bit mixed precision
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flux
seed: 42
walk_seed: true
guidance_scale: 1 # schnell does not do guidance
sample_steps: 4 # 1 - 4 works well
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,107 @@
---
# This configuration requires 48GB of VRAM or more to operate
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flex_finetune_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
save:
dtype: bf16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 2 # how many intermittent saves to keep
save_format: 'diffusers' # 'diffusers'
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
# cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
train:
batch_size: 1
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
bypass_guidance_embedding: true
# can be 'sigmoid', 'linear', or 'lognorm_blend'
timestep_type: 'sigmoid'
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with flex
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adafactor"
lr: 3e-5
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
# do_paramiter_swapping: true
# paramiter_swapping_factor: 0.9
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "ostris/Flex.1-alpha"
is_flux: true # flex is flux architecture
# full finetuning quantized models is a crapshoot and results in subpar outputs
# quantize: true
# you can quantize just the T5 text encoder here to save vram
quantize_te: true
# only train the transformer blocks
only_if_contains:
- "transformer.transformer_blocks."
- "transformer.single_transformer_blocks."
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flex
seed: 42
walk_seed: true
guidance_scale: 4
sample_steps: 25
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,99 @@
---
# This configuration requires 24GB of VRAM or more to operate
job: extension
config:
# this name will be the folder and filename name
name: "my_first_lumina_finetune_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
save:
dtype: bf16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 2 # how many intermittent saves to keep
save_format: 'diffusers' # 'diffusers'
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
# cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
train:
batch_size: 1
# can be 'sigmoid', 'linear', or 'lumina2_shift'
timestep_type: 'lumina2_shift'
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with lumina2
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adafactor"
lr: 3e-5
# Paramiter swapping can reduce vram requirements. Set factor from 1.0 to 0.0.
# 0.1 is 10% of paramiters active at easc step. Only works with adafactor
# do_paramiter_swapping: true
# paramiter_swapping_factor: 0.9
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
# ema_config:
# use_ema: true
# ema_decay: 0.99
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
is_lumina2: true # lumina2 architecture
# you can quantize just the Gemma2 text encoder here to save vram
quantize_te: true
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: ""
seed: 42
walk_seed: true
guidance_scale: 4.0
sample_steps: 25
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,101 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flex_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flex enjoys multiple resolutions
train:
batch_size: 1
# IMPORTANT! For Flex, you must bypass the guidance embedder during training
bypass_guidance_embedding: true
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with flex
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new vell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flex, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "ostris/Flex.1-alpha"
is_flux: true
quantize: true # run 8bit mixed precision
quantize_kwargs:
exclude:
- "*time_text_embed*" # exclude the time text embedder from quantization
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flex
seed: 42
walk_seed: true
guidance_scale: 4
sample_steps: 25
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,96 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flux_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation_steps: 1
train_unet: true
train_text_encoder: false # probably won't work with flux
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new vell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "black-forest-labs/FLUX.1-dev"
is_flux: true
quantize: true # run 8bit mixed precision
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flux
seed: 42
walk_seed: true
guidance_scale: 4
sample_steps: 20
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,98 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_flux_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # flux enjoys multiple resolutions
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation_steps: 1
train_unet: true
train_text_encoder: false # probably won't work with flux
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new bell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for flux, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "black-forest-labs/FLUX.1-schnell"
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter" # Required for flux schnell training
is_flux: true
quantize: true # run 8bit mixed precision
# low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary
# low_vram: true # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: "" # not used on flux
seed: 42
walk_seed: true
guidance_scale: 1 # schnell does not do guidance
sample_steps: 4 # 1 - 4 works well
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,96 @@
---
# This configuration requires 20GB of VRAM or more to operate
job: extension
config:
# this name will be the folder and filename name
name: "my_first_lumina_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: bf16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 2 # how many intermittent saves to keep
save_format: 'diffusers' # 'diffusers'
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
# cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 512, 768, 1024 ] # lumina2 enjoys multiple resolutions
train:
batch_size: 1
# can be 'sigmoid', 'linear', or 'lumina2_shift'
timestep_type: 'lumina2_shift'
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with lumina2
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# ema will smooth out learning, but could slow it down. Recommended to leave on if you have the vram
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for lumina2, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "Alpha-VLLM/Lumina-Image-2.0"
is_lumina2: true # lumina2 architecture
# you can quantize just the Gemma2 text encoder here to save vram
quantize_te: true
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a cat that is half black and half orange tabby, split down the middle. The cat has on a blue tophat. They are holding a martini glass with a pink ball of yarn in it with green knitting needles sticking out, in one paw. In the other paw, they are holding a DVD case for a movie titled, \"This is a test\" that has a golden robot on it. In the background is a busy night club with a giant mushroom man dancing with a bear."
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: ""
seed: 42
walk_seed: true
guidance_scale: 4.0
sample_steps: 25
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,97 @@
---
# NOTE!! THIS IS CURRENTLY EXPERIMENTAL AND UNDER DEVELOPMENT. SOME THINGS WILL CHANGE
job: extension
config:
# this name will be the folder and filename name
name: "my_first_sd3l_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 16
linear_alpha: 16
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 1024 ]
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation_steps: 1
train_unet: true
train_text_encoder: false # May not fully work with SD3 yet
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch"
timestep_type: "linear" # linear or sigmoid
optimizer: "adamw8bit"
lr: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# uncomment to use new vell curved weighting. Experimental but may produce better results
# linear_timesteps: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
# will probably need this if gpu supports it for sd3, other dtypes may not work correctly
dtype: bf16
model:
# huggingface model name or path
name_or_path: "stabilityai/stable-diffusion-3.5-large"
is_v3: true
quantize: true # run 8bit mixed precision
sample:
sampler: "flowmatch" # must match train.noise_scheduler
sample_every: 250 # sample every this many steps
width: 1024
height: 1024
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman with red hair, playing chess at the park, bomb going off in the background"
- "a woman holding a coffee cup, in a beanie, sitting at a cafe"
- "a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini"
- "a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background"
- "a bear building a log cabin in the snow covered mountains"
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
- "hipster man with a beard, building a chair, in a wood shop"
- "photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop"
- "a man holding a sign that says, 'this is a sign'"
- "a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle"
neg: ""
seed: 42
walk_seed: true
guidance_scale: 4
sample_steps: 25
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,101 @@
# IMPORTANT: The Wan2.1 14B model is huge. This config should work on 24GB GPUs. It cannot
# support keeping the text encoder on GPU while training with 24GB, so it is only good
# for training on a single prompt, for example a person with a trigger word.
# to train on captions, you need more vran for now.
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_wan21_14b_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# this is probably needed for 24GB cards when offloading TE to CPU
trigger_word: "p3r5on"
network:
type: "lora"
linear: 32
linear_alpha: 32
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
# it works well for characters, but not as well for "actions"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 632 ] # will be around 480p
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with wan
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
timestep_type: 'sigmoid'
optimizer: "adamw8bit"
lr: 1e-4
optimizer_params:
weight_decay: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
dtype: bf16
# required for 24GB cards
# this will encode your trigger word and use those embeddings for every image in the dataset
unload_text_encoder: true
model:
# huggingface model name or path
name_or_path: "Wan-AI/Wan2.1-T2V-14B-Diffusers"
arch: 'wan21'
# these settings will save as much vram as possible
quantize: true
quantize_te: true
low_vram: true
sample:
sampler: "flowmatch"
sample_every: 250 # sample every this many steps
width: 832
height: 480
num_frames: 40
fps: 15
# samples take a long time. so use them sparingly
# samples will be animated webp files, if you don't see them animated, open in a browser.
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
neg: ""
seed: 42
walk_seed: true
guidance_scale: 5
sample_steps: 30
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

View File

@@ -0,0 +1,90 @@
---
job: extension
config:
# this name will be the folder and filename name
name: "my_first_wan21_1b_lora_v1"
process:
- type: 'sd_trainer'
# root folder to save training sessions/samples/weights
training_folder: "output"
# uncomment to see performance stats in the terminal every N steps
# performance_log_every: 1000
device: cuda:0
# if a trigger word is specified, it will be added to captions of training data if it does not already exist
# alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word
# trigger_word: "p3r5on"
network:
type: "lora"
linear: 32
linear_alpha: 32
save:
dtype: float16 # precision to save
save_every: 250 # save every this many steps
max_step_saves_to_keep: 4 # how many intermittent saves to keep
push_to_hub: false #change this to True to push your trained model to Hugging Face.
# You can either set up a HF_TOKEN env variable or you'll be prompted to log-in
# hf_repo_id: your-username/your-model-slug
# hf_private: true #whether the repo is private or public
datasets:
# datasets are a folder of images. captions need to be txt files with the same name as the image
# for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently
# images will automatically be resized and bucketed into the resolution specified
# on windows, escape back slashes with another backslash so
# "C:\\path\\to\\images\\folder"
# AI-Toolkit does not currently support video datasets, we will train on 1 frame at a time
# it works well for characters, but not as well for "actions"
- folder_path: "/path/to/images/folder"
caption_ext: "txt"
caption_dropout_rate: 0.05 # will drop out the caption 5% of time
shuffle_tokens: false # shuffle caption order, split by commas
cache_latents_to_disk: true # leave this true unless you know what you're doing
resolution: [ 632 ] # will be around 480p
train:
batch_size: 1
steps: 2000 # total number of steps to train 500 - 4000 is a good range
gradient_accumulation: 1
train_unet: true
train_text_encoder: false # probably won't work with wan
gradient_checkpointing: true # need the on unless you have a ton of vram
noise_scheduler: "flowmatch" # for training only
timestep_type: 'sigmoid'
optimizer: "adamw8bit"
lr: 1e-4
optimizer_params:
weight_decay: 1e-4
# uncomment this to skip the pre training sample
# skip_first_sample: true
# uncomment to completely disable sampling
# disable_sampling: true
# ema will smooth out learning, but could slow it down. Recommended to leave on.
ema_config:
use_ema: true
ema_decay: 0.99
dtype: bf16
model:
# huggingface model name or path
name_or_path: "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
arch: 'wan21'
quantize_te: true # saves vram
sample:
sampler: "flowmatch"
sample_every: 250 # sample every this many steps
width: 832
height: 480
num_frames: 40
fps: 15
# samples take a long time. so use them sparingly
# samples will be animated webp files, if you don't see them animated, open in a browser.
prompts:
# you can add [trigger] to the prompts here and it will be replaced with the trigger word
# - "[trigger] holding a sign that says 'I LOVE PROMPTS!'"\
- "woman playing the guitar, on stage, singing a song, laser lights, punk rocker"
neg: ""
seed: 42
walk_seed: true
guidance_scale: 5
sample_steps: 30
# you can add any additional meta info here. [name] is replaced with config name at top
meta:
name: "[name]"
version: '1.0'

25
docker-compose.yml Normal file
View File

@@ -0,0 +1,25 @@
version: "3.8"
services:
ai-toolkit:
image: ostris/aitoolkit:latest
restart: unless-stopped
ports:
- "8675:8675"
volumes:
- ~/.cache/huggingface/hub:/root/.cache/huggingface/hub
- ./aitk_db.db:/app/ai-toolkit/aitk_db.db
- ./datasets:/app/ai-toolkit/datasets
- ./output:/app/ai-toolkit/output
- ./config:/app/ai-toolkit/config
environment:
- AI_TOOLKIT_AUTH=${AI_TOOLKIT_AUTH:-password}
- NODE_ENV=production
- TZ=UTC
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: all
capabilities: [gpu]

67
docker/Dockerfile Normal file
View File

@@ -0,0 +1,67 @@
FROM nvidia/cuda:12.6.3-base-ubuntu22.04
LABEL authors="jaret"
# Set noninteractive to avoid timezone prompts
ENV DEBIAN_FRONTEND=noninteractive
# Install dependencies
RUN apt-get update && apt-get install --no-install-recommends -y \
git \
curl \
build-essential \
cmake \
wget \
python3.10 \
python3-pip \
python3-dev \
python3-setuptools \
python3-wheel \
python3-venv \
ffmpeg \
tmux \
htop \
nvtop \
python3-opencv \
&& apt-get clean \
&& rm -rf /var/lib/apt/lists/*
# Install nodejs
WORKDIR /tmp
RUN curl -sL https://deb.nodesource.com/setup_23.x -o nodesource_setup.sh && \
bash nodesource_setup.sh && \
apt-get update && \
apt-get install -y nodejs && \
apt-get clean && \
rm -rf /var/lib/apt/lists/*
WORKDIR /app
# Set aliases for python and pip
RUN ln -s /usr/bin/python3 /usr/bin/python
# install pytorch before cache bust to avoid redownloading pytorch
RUN pip install --no-cache-dir torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/cu126
# Fix cache busting by moving CACHEBUST to right before git clone
ARG CACHEBUST=1234
RUN echo "Cache bust: ${CACHEBUST}" && \
git clone https://github.com/ostris/ai-toolkit.git && \
cd ai-toolkit && \
git submodule update --init --recursive
WORKDIR /app/ai-toolkit
# Install Python dependencies
RUN pip install --no-cache-dir -r requirements.txt
# Build UI
WORKDIR /app/ai-toolkit/ui
RUN npm install && \
npm run build && \
npm run update_db
# Expose port (assuming the application runs on port 3000)
EXPOSE 8675
CMD ["npm", "run", "start"]

View File

@@ -0,0 +1,256 @@
import math
import os
import random
from collections import OrderedDict
from typing import List
import numpy as np
from PIL import Image
from diffusers import T2IAdapter
from diffusers.utils.torch_utils import randn_tensor
from torch.utils.data import DataLoader
from diffusers import StableDiffusionXLImg2ImgPipeline, PixArtSigmaPipeline
from tqdm import tqdm
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
from toolkit.sampler import get_sampler
from toolkit.stable_diffusion_model import StableDiffusion
import gc
import torch
from jobs.process import BaseExtensionProcess
from toolkit.data_loader import get_dataloader_from_datasets
from toolkit.train_tools import get_torch_dtype
from controlnet_aux.midas import MidasDetector
from diffusers.utils import load_image
from torchvision.transforms import ToTensor
def flush():
torch.cuda.empty_cache()
gc.collect()
class GenerateConfig:
def __init__(self, **kwargs):
self.prompts: List[str]
self.sampler = kwargs.get('sampler', 'ddpm')
self.neg = kwargs.get('neg', '')
self.seed = kwargs.get('seed', -1)
self.walk_seed = kwargs.get('walk_seed', False)
self.guidance_scale = kwargs.get('guidance_scale', 7)
self.sample_steps = kwargs.get('sample_steps', 20)
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
self.ext = kwargs.get('ext', 'png')
self.denoise_strength = kwargs.get('denoise_strength', 0.5)
self.trigger_word = kwargs.get('trigger_word', None)
class Img2ImgGenerator(BaseExtensionProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.output_folder = self.get_conf('output_folder', required=True)
self.copy_inputs_to = self.get_conf('copy_inputs_to', None)
self.device = self.get_conf('device', 'cuda')
self.model_config = ModelConfig(**self.get_conf('model', required=True))
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
self.is_latents_cached = True
raw_datasets = self.get_conf('datasets', None)
if raw_datasets is not None and len(raw_datasets) > 0:
raw_datasets = preprocess_dataset_raw_config(raw_datasets)
self.datasets = None
self.datasets_reg = None
self.dtype = self.get_conf('dtype', 'float16')
self.torch_dtype = get_torch_dtype(self.dtype)
self.params = []
if raw_datasets is not None and len(raw_datasets) > 0:
for raw_dataset in raw_datasets:
dataset = DatasetConfig(**raw_dataset)
is_caching = dataset.cache_latents or dataset.cache_latents_to_disk
if not is_caching:
self.is_latents_cached = False
if dataset.is_reg:
if self.datasets_reg is None:
self.datasets_reg = []
self.datasets_reg.append(dataset)
else:
if self.datasets is None:
self.datasets = []
self.datasets.append(dataset)
self.progress_bar = None
self.sd = StableDiffusion(
device=self.device,
model_config=self.model_config,
dtype=self.dtype,
)
print(f"Using device {self.device}")
self.data_loader: DataLoader = None
self.adapter: T2IAdapter = None
def to_pil(self, img):
# image comes in -1 to 1. convert to a PIL RGB image
img = (img + 1) / 2
img = img.clamp(0, 1)
img = img[0].permute(1, 2, 0).cpu().numpy()
img = (img * 255).astype(np.uint8)
image = Image.fromarray(img)
return image
def run(self):
with torch.no_grad():
super().run()
print("Loading model...")
self.sd.load_model()
device = torch.device(self.device)
if self.model_config.is_xl:
pipe = StableDiffusionXLImg2ImgPipeline(
vae=self.sd.vae,
unet=self.sd.unet,
text_encoder=self.sd.text_encoder[0],
text_encoder_2=self.sd.text_encoder[1],
tokenizer=self.sd.tokenizer[0],
tokenizer_2=self.sd.tokenizer[1],
scheduler=get_sampler(self.generate_config.sampler),
).to(device, dtype=self.torch_dtype)
elif self.model_config.is_pixart:
pipe = self.sd.pipeline.to(device, dtype=self.torch_dtype)
else:
raise NotImplementedError("Only XL models are supported")
pipe.set_progress_bar_config(disable=True)
# pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
num_batches = len(self.data_loader)
pbar = tqdm(total=num_batches, desc="Generating images")
seed = self.generate_config.seed
# load images from datasets, use tqdm
for i, batch in enumerate(self.data_loader):
batch: DataLoaderBatchDTO = batch
gen_seed = seed if seed > 0 else random.randint(0, 2 ** 32 - 1)
generator = torch.manual_seed(gen_seed)
file_item: FileItemDTO = batch.file_items[0]
img_path = file_item.path
img_filename = os.path.basename(img_path)
img_filename_no_ext = os.path.splitext(img_filename)[0]
img_filename = img_filename_no_ext + '.' + self.generate_config.ext
output_path = os.path.join(self.output_folder, img_filename)
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
if self.copy_inputs_to is not None:
output_inputs_path = os.path.join(self.copy_inputs_to, img_filename)
output_inputs_caption_path = os.path.join(self.copy_inputs_to, img_filename_no_ext + '.txt')
else:
output_inputs_path = None
output_inputs_caption_path = None
caption = batch.get_caption_list()[0]
if self.generate_config.trigger_word is not None:
caption = caption.replace('[trigger]', self.generate_config.trigger_word)
img: torch.Tensor = batch.tensor.clone()
image = self.to_pil(img)
# image.save(output_depth_path)
if self.model_config.is_pixart:
pipe: PixArtSigmaPipeline = pipe
# Encode the full image once
encoded_image = pipe.vae.encode(
pipe.image_processor.preprocess(image).to(device=pipe.device, dtype=pipe.dtype))
if hasattr(encoded_image, "latent_dist"):
latents = encoded_image.latent_dist.sample(generator)
elif hasattr(encoded_image, "latents"):
latents = encoded_image.latents
else:
raise AttributeError("Could not access latents of provided encoder_output")
latents = pipe.vae.config.scaling_factor * latents
# latents = self.sd.encode_images(img)
# self.sd.noise_scheduler.set_timesteps(self.generate_config.sample_steps)
# start_step = math.floor(self.generate_config.sample_steps * self.generate_config.denoise_strength)
# timestep = self.sd.noise_scheduler.timesteps[start_step].unsqueeze(0)
# timestep = timestep.to(device, dtype=torch.int32)
# latent = latent.to(device, dtype=self.torch_dtype)
# noise = torch.randn_like(latent, device=device, dtype=self.torch_dtype)
# latent = self.sd.add_noise(latent, noise, timestep)
# timesteps_to_use = self.sd.noise_scheduler.timesteps[start_step + 1:]
batch_size = 1
num_images_per_prompt = 1
shape = (batch_size, pipe.transformer.config.in_channels, image.height // pipe.vae_scale_factor,
image.width // pipe.vae_scale_factor)
noise = randn_tensor(shape, generator=generator, device=pipe.device, dtype=pipe.dtype)
# noise = torch.randn_like(latents, device=device, dtype=self.torch_dtype)
num_inference_steps = self.generate_config.sample_steps
strength = self.generate_config.denoise_strength
# Get timesteps
init_timestep = min(int(num_inference_steps * strength), num_inference_steps)
t_start = max(num_inference_steps - init_timestep, 0)
pipe.scheduler.set_timesteps(num_inference_steps, device="cpu")
timesteps = pipe.scheduler.timesteps[t_start:]
timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
latents = pipe.scheduler.add_noise(latents, noise, timestep)
gen_images = pipe.__call__(
prompt=caption,
negative_prompt=self.generate_config.neg,
latents=latents,
timesteps=timesteps,
width=image.width,
height=image.height,
num_inference_steps=num_inference_steps,
num_images_per_prompt=num_images_per_prompt,
guidance_scale=self.generate_config.guidance_scale,
# strength=self.generate_config.denoise_strength,
use_resolution_binning=False,
output_type="np"
).images[0]
gen_images = (gen_images * 255).clip(0, 255).astype(np.uint8)
gen_images = Image.fromarray(gen_images)
else:
pipe: StableDiffusionXLImg2ImgPipeline = pipe
gen_images = pipe.__call__(
prompt=caption,
negative_prompt=self.generate_config.neg,
image=image,
num_inference_steps=self.generate_config.sample_steps,
guidance_scale=self.generate_config.guidance_scale,
strength=self.generate_config.denoise_strength,
).images[0]
os.makedirs(os.path.dirname(output_path), exist_ok=True)
gen_images.save(output_path)
# save caption
with open(output_caption_path, 'w') as f:
f.write(caption)
if output_inputs_path is not None:
os.makedirs(os.path.dirname(output_inputs_path), exist_ok=True)
image.save(output_inputs_path)
with open(output_inputs_caption_path, 'w') as f:
f.write(caption)
pbar.update(1)
batch.cleanup()
pbar.close()
print("Done generating images")
# cleanup
del self.sd
gc.collect()
torch.cuda.empty_cache()

View File

@@ -7,7 +7,7 @@ import numpy as np
from PIL import Image
from diffusers import T2IAdapter
from torch.utils.data import DataLoader
from diffusers import StableDiffusionXLAdapterPipeline
from diffusers import StableDiffusionXLAdapterPipeline, StableDiffusionAdapterPipeline
from tqdm import tqdm
from toolkit.config_modules import ModelConfig, GenerateImageConfig, preprocess_dataset_raw_config, DatasetConfig
@@ -100,25 +100,43 @@ class ReferenceGenerator(BaseExtensionProcess):
if self.generate_config.t2i_adapter_path is not None:
self.adapter = T2IAdapter.from_pretrained(
"TencentARC/t2i-adapter-depth-midas-sdxl-1.0", torch_dtype=self.torch_dtype, varient="fp16"
self.generate_config.t2i_adapter_path,
torch_dtype=self.torch_dtype,
varient="fp16"
).to(device)
midas_depth = MidasDetector.from_pretrained(
"valhalla/t2iadapter-aux-models", filename="dpt_large_384.pt", model_type="dpt_large"
).to(device)
pipe = StableDiffusionXLAdapterPipeline(
vae=self.sd.vae,
unet=self.sd.unet,
text_encoder=self.sd.text_encoder[0],
text_encoder_2=self.sd.text_encoder[1],
tokenizer=self.sd.tokenizer[0],
tokenizer_2=self.sd.tokenizer[1],
scheduler=get_sampler(self.generate_config.sampler),
adapter=self.adapter,
).to(device)
if self.model_config.is_xl:
pipe = StableDiffusionXLAdapterPipeline(
vae=self.sd.vae,
unet=self.sd.unet,
text_encoder=self.sd.text_encoder[0],
text_encoder_2=self.sd.text_encoder[1],
tokenizer=self.sd.tokenizer[0],
tokenizer_2=self.sd.tokenizer[1],
scheduler=get_sampler(self.generate_config.sampler),
adapter=self.adapter,
).to(device, dtype=self.torch_dtype)
else:
pipe = StableDiffusionAdapterPipeline(
vae=self.sd.vae,
unet=self.sd.unet,
text_encoder=self.sd.text_encoder,
tokenizer=self.sd.tokenizer,
scheduler=get_sampler(self.generate_config.sampler),
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
adapter=self.adapter,
).to(device, dtype=self.torch_dtype)
pipe.set_progress_bar_config(disable=True)
pipe.unet = torch.compile(pipe.unet, mode="reduce-overhead", fullgraph=True)
# midas_depth = torch.compile(midas_depth, mode="reduce-overhead", fullgraph=True)
self.data_loader = get_dataloader_from_datasets(self.datasets, 1, self.sd)
num_batches = len(self.data_loader)
@@ -176,6 +194,7 @@ class ReferenceGenerator(BaseExtensionProcess):
adapter_conditioning_scale=self.generate_config.adapter_conditioning_scale,
guidance_scale=self.generate_config.guidance_scale,
).images[0]
os.makedirs(os.path.dirname(output_path), exist_ok=True)
gen_images.save(output_path)
# save caption

View File

@@ -36,7 +36,24 @@ class PureLoraGenerator(Extension):
return PureLoraGenerator
# This is for generic training (LoRA, Dreambooth, FineTuning)
class Img2ImgGeneratorExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "batch_img2img"
# name is the name of the extension for printing
name = "Img2ImgGeneratorExtension"
# This is where your process class is loaded
# keep your imports in here so they don't slow down the rest of the program
@classmethod
def get_process(cls):
# import your process class here so it is only loaded when needed and return it
from .Img2ImgGenerator import Img2ImgGenerator
return Img2ImgGenerator
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
AdvancedReferenceGeneratorExtension, PureLoraGenerator
AdvancedReferenceGeneratorExtension, PureLoraGenerator, Img2ImgGeneratorExtension
]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,234 @@
from collections import OrderedDict
import os
import sqlite3
import asyncio
import concurrent.futures
from extensions_built_in.sd_trainer.SDTrainer import SDTrainer
from typing import Literal, Optional
AITK_Status = Literal["running", "stopped", "error", "completed"]
class UITrainer(SDTrainer):
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super(UITrainer, self).__init__(process_id, job, config, **kwargs)
self.sqlite_db_path = self.config.get("sqlite_db_path", "./aitk_db.db")
if not os.path.exists(self.sqlite_db_path):
raise Exception(
f"SQLite database not found at {self.sqlite_db_path}")
print(f"Using SQLite database at {self.sqlite_db_path}")
self.job_id = os.environ.get("AITK_JOB_ID", None)
self.job_id = self.job_id.strip() if self.job_id is not None else None
print(f"Job ID: \"{self.job_id}\"")
if self.job_id is None:
raise Exception("AITK_JOB_ID not set")
self.is_stopping = False
# Create a thread pool for database operations
self.thread_pool = concurrent.futures.ThreadPoolExecutor(max_workers=1)
# Track all async tasks
self._async_tasks = []
# Initialize the status
self._run_async_operation(self._update_status("running", "Starting"))
def _run_async_operation(self, coro):
"""Helper method to run an async coroutine and track the task."""
try:
loop = asyncio.get_event_loop()
except RuntimeError:
# No event loop exists, create a new one
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# Create a task and track it
if loop.is_running():
task = asyncio.run_coroutine_threadsafe(coro, loop)
self._async_tasks.append(asyncio.wrap_future(task))
else:
task = loop.create_task(coro)
self._async_tasks.append(task)
loop.run_until_complete(task)
async def _execute_db_operation(self, operation_func):
"""Execute a database operation in a separate thread to avoid blocking."""
loop = asyncio.get_event_loop()
return await loop.run_in_executor(self.thread_pool, operation_func)
def _db_connect(self):
"""Create a new connection for each operation to avoid locking."""
conn = sqlite3.connect(self.sqlite_db_path, timeout=10.0)
conn.isolation_level = None # Enable autocommit mode
return conn
def should_stop(self):
def _check_stop():
with self._db_connect() as conn:
cursor = conn.cursor()
cursor.execute(
"SELECT stop FROM Job WHERE id = ?", (self.job_id,))
stop = cursor.fetchone()
return False if stop is None else stop[0] == 1
return _check_stop()
def maybe_stop(self):
if self.should_stop():
self._run_async_operation(
self._update_status("stopped", "Job stopped"))
self.is_stopping = True
raise Exception("Job stopped")
async def _update_key(self, key, value):
if not self.accelerator.is_main_process:
return
def _do_update():
with self._db_connect() as conn:
cursor = conn.cursor()
cursor.execute("BEGIN IMMEDIATE")
try:
# Convert the value to string if it's not already
if isinstance(value, str):
value_to_insert = value
else:
value_to_insert = str(value)
# Use parameterized query for both the column name and value
update_query = f"UPDATE Job SET {key} = ? WHERE id = ?"
cursor.execute(
update_query, (value_to_insert, self.job_id))
finally:
cursor.execute("COMMIT")
await self._execute_db_operation(_do_update)
def update_step(self):
"""Non-blocking update of the step count."""
if self.accelerator.is_main_process:
self._run_async_operation(self._update_key("step", self.step_num))
def update_db_key(self, key, value):
"""Non-blocking update a key in the database."""
if self.accelerator.is_main_process:
self._run_async_operation(self._update_key(key, value))
async def _update_status(self, status: AITK_Status, info: Optional[str] = None):
if not self.accelerator.is_main_process:
return
def _do_update():
with self._db_connect() as conn:
cursor = conn.cursor()
cursor.execute("BEGIN IMMEDIATE")
try:
if info is not None:
cursor.execute(
"UPDATE Job SET status = ?, info = ? WHERE id = ?",
(status, info, self.job_id)
)
else:
cursor.execute(
"UPDATE Job SET status = ? WHERE id = ?",
(status, self.job_id)
)
finally:
cursor.execute("COMMIT")
await self._execute_db_operation(_do_update)
def update_status(self, status: AITK_Status, info: Optional[str] = None):
"""Non-blocking update of status."""
if self.accelerator.is_main_process:
self._run_async_operation(self._update_status(status, info))
async def wait_for_all_async(self):
"""Wait for all tracked async operations to complete."""
if not self._async_tasks:
return
try:
await asyncio.gather(*self._async_tasks)
except Exception as e:
pass
finally:
# Clear the task list after completion
self._async_tasks.clear()
def on_error(self, e: Exception):
super(UITrainer, self).on_error(e)
if self.accelerator.is_main_process and not self.is_stopping:
self.update_status("error", str(e))
self.update_db_key("step", self.last_save_step)
asyncio.run(self.wait_for_all_async())
self.thread_pool.shutdown(wait=True)
def handle_timing_print_hook(self, timing_dict):
if "train_loop" not in timing_dict:
print("train_loop not found in timing_dict", timing_dict)
return
seconds_per_iter = timing_dict["train_loop"]
# determine iter/sec or sec/iter
if seconds_per_iter < 1:
iters_per_sec = 1 / seconds_per_iter
self.update_db_key("speed_string", f"{iters_per_sec:.2f} iter/sec")
else:
self.update_db_key(
"speed_string", f"{seconds_per_iter:.2f} sec/iter")
def done_hook(self):
super(UITrainer, self).done_hook()
self.update_status("completed", "Training completed")
# Wait for all async operations to finish before shutting down
asyncio.run(self.wait_for_all_async())
self.thread_pool.shutdown(wait=True)
def end_step_hook(self):
super(UITrainer, self).end_step_hook()
self.update_step()
self.maybe_stop()
def hook_before_model_load(self):
super().hook_before_model_load()
self.maybe_stop()
self.update_status("running", "Loading model")
def before_dataset_load(self):
super().before_dataset_load()
self.maybe_stop()
self.update_status("running", "Loading dataset")
def hook_before_train_loop(self):
super().hook_before_train_loop()
self.maybe_stop()
self.update_step()
self.update_status("running", "Training")
self.timer.add_after_print_hook(self.handle_timing_print_hook)
def status_update_hook_func(self, string):
self.update_status("running", string)
def hook_after_sd_init_before_load(self):
super().hook_after_sd_init_before_load()
self.maybe_stop()
self.sd.add_status_update_hook(self.status_update_hook_func)
def sample_step_hook(self, img_num, total_imgs):
super().sample_step_hook(img_num, total_imgs)
self.maybe_stop()
self.update_status(
"running", f"Generating images - {img_num + 1}/{total_imgs}")
def sample(self, step=None, is_first=False):
self.maybe_stop()
total_imgs = len(self.sample_config.prompts)
self.update_status("running", f"Generating images - 0/{total_imgs}")
super().sample(step, is_first)
self.maybe_stop()
self.update_status("running", "Training")
def save(self, step=None):
self.maybe_stop()
self.update_status("running", "Saving model")
super().save(step)
self.maybe_stop()
self.update_status("running", "Training")

View File

@@ -18,6 +18,22 @@ class SDTrainerExtension(Extension):
from .SDTrainer import SDTrainer
return SDTrainer
# This is for generic training (LoRA, Dreambooth, FineTuning)
class UITrainerExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "ui_trainer"
# name is the name of the extension for printing
name = "UI Trainer"
# This is where your process class is loaded
# keep your imports in here so they don't slow down the rest of the program
@classmethod
def get_process(cls):
# import your process class here so it is only loaded when needed and return it
from .UITrainer import UITrainer
return UITrainer
# for backwards compatability
class TextualInversionTrainer(SDTrainerExtension):
@@ -26,5 +42,5 @@ class TextualInversionTrainer(SDTrainerExtension):
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
SDTrainerExtension, TextualInversionTrainer
SDTrainerExtension, TextualInversionTrainer, UITrainerExtension
]

414
flux_train_ui.py Normal file
View File

@@ -0,0 +1,414 @@
import os
from huggingface_hub import whoami
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
import sys
# Add the current working directory to the Python path
sys.path.insert(0, os.getcwd())
import gradio as gr
from PIL import Image
import torch
import uuid
import os
import shutil
import json
import yaml
from slugify import slugify
from transformers import AutoProcessor, AutoModelForCausalLM
sys.path.insert(0, "ai-toolkit")
from toolkit.job import get_job
MAX_IMAGES = 150
def load_captioning(uploaded_files, concept_sentence):
uploaded_images = [file for file in uploaded_files if not file.endswith('.txt')]
txt_files = [file for file in uploaded_files if file.endswith('.txt')]
txt_files_dict = {os.path.splitext(os.path.basename(txt_file))[0]: txt_file for txt_file in txt_files}
updates = []
if len(uploaded_images) <= 1:
raise gr.Error(
"Please upload at least 2 images to train your model (the ideal number with default settings is between 4-30)"
)
elif len(uploaded_images) > MAX_IMAGES:
raise gr.Error(f"For now, only {MAX_IMAGES} or less images are allowed for training")
# Update for the captioning_area
# for _ in range(3):
updates.append(gr.update(visible=True))
# Update visibility and image for each captioning row and image
for i in range(1, MAX_IMAGES + 1):
# Determine if the current row and image should be visible
visible = i <= len(uploaded_images)
# Update visibility of the captioning row
updates.append(gr.update(visible=visible))
# Update for image component - display image if available, otherwise hide
image_value = uploaded_images[i - 1] if visible else None
updates.append(gr.update(value=image_value, visible=visible))
corresponding_caption = False
if(image_value):
base_name = os.path.splitext(os.path.basename(image_value))[0]
print(base_name)
print(image_value)
if base_name in txt_files_dict:
print("entrou")
with open(txt_files_dict[base_name], 'r') as file:
corresponding_caption = file.read()
# Update value of captioning area
text_value = corresponding_caption if visible and corresponding_caption else "[trigger]" if visible and concept_sentence else None
updates.append(gr.update(value=text_value, visible=visible))
# Update for the sample caption area
updates.append(gr.update(visible=True))
# Update prompt samples
updates.append(gr.update(placeholder=f'A portrait of person in a bustling cafe {concept_sentence}', value=f'A person in a bustling cafe {concept_sentence}'))
updates.append(gr.update(placeholder=f"A mountainous landscape in the style of {concept_sentence}"))
updates.append(gr.update(placeholder=f"A {concept_sentence} in a mall"))
updates.append(gr.update(visible=True))
return updates
def hide_captioning():
return gr.update(visible=False), gr.update(visible=False), gr.update(visible=False)
def create_dataset(*inputs):
print("Creating dataset")
images = inputs[0]
destination_folder = str(f"datasets/{uuid.uuid4()}")
if not os.path.exists(destination_folder):
os.makedirs(destination_folder)
jsonl_file_path = os.path.join(destination_folder, "metadata.jsonl")
with open(jsonl_file_path, "a") as jsonl_file:
for index, image in enumerate(images):
new_image_path = shutil.copy(image, destination_folder)
original_caption = inputs[index + 1]
file_name = os.path.basename(new_image_path)
data = {"file_name": file_name, "prompt": original_caption}
jsonl_file.write(json.dumps(data) + "\n")
return destination_folder
def run_captioning(images, concept_sentence, *captions):
#Load internally to not consume resources for training
device = "cuda" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16
model = AutoModelForCausalLM.from_pretrained(
"multimodalart/Florence-2-large-no-flash-attn", torch_dtype=torch_dtype, trust_remote_code=True
).to(device)
processor = AutoProcessor.from_pretrained("multimodalart/Florence-2-large-no-flash-attn", trust_remote_code=True)
captions = list(captions)
for i, image_path in enumerate(images):
print(captions[i])
if isinstance(image_path, str): # If image is a file path
image = Image.open(image_path).convert("RGB")
prompt = "<DETAILED_CAPTION>"
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
generated_ids = model.generate(
input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=1024, num_beams=3
)
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
parsed_answer = processor.post_process_generation(
generated_text, task=prompt, image_size=(image.width, image.height)
)
caption_text = parsed_answer["<DETAILED_CAPTION>"].replace("The image shows ", "")
if concept_sentence:
caption_text = f"{caption_text} [trigger]"
captions[i] = caption_text
yield captions
model.to("cpu")
del model
del processor
def recursive_update(d, u):
for k, v in u.items():
if isinstance(v, dict) and v:
d[k] = recursive_update(d.get(k, {}), v)
else:
d[k] = v
return d
def start_training(
lora_name,
concept_sentence,
steps,
lr,
rank,
model_to_train,
low_vram,
dataset_folder,
sample_1,
sample_2,
sample_3,
use_more_advanced_options,
more_advanced_options,
):
push_to_hub = True
if not lora_name:
raise gr.Error("You forgot to insert your LoRA name! This name has to be unique.")
try:
if whoami()["auth"]["accessToken"]["role"] == "write" or "repo.write" in whoami()["auth"]["accessToken"]["fineGrained"]["scoped"][0]["permissions"]:
gr.Info(f"Starting training locally {whoami()['name']}. Your LoRA will be available locally and in Hugging Face after it finishes.")
else:
push_to_hub = False
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
except:
push_to_hub = False
gr.Warning("Started training locally. Your LoRa will only be available locally because you didn't login with a `write` token to Hugging Face")
print("Started training")
slugged_lora_name = slugify(lora_name)
# Load the default config
with open("config/examples/train_lora_flux_24gb.yaml", "r") as f:
config = yaml.safe_load(f)
# Update the config with user inputs
config["config"]["name"] = slugged_lora_name
config["config"]["process"][0]["model"]["low_vram"] = low_vram
config["config"]["process"][0]["train"]["skip_first_sample"] = True
config["config"]["process"][0]["train"]["steps"] = int(steps)
config["config"]["process"][0]["train"]["lr"] = float(lr)
config["config"]["process"][0]["network"]["linear"] = int(rank)
config["config"]["process"][0]["network"]["linear_alpha"] = int(rank)
config["config"]["process"][0]["datasets"][0]["folder_path"] = dataset_folder
config["config"]["process"][0]["save"]["push_to_hub"] = push_to_hub
if(push_to_hub):
try:
username = whoami()["name"]
except:
raise gr.Error("Error trying to retrieve your username. Are you sure you are logged in with Hugging Face?")
config["config"]["process"][0]["save"]["hf_repo_id"] = f"{username}/{slugged_lora_name}"
config["config"]["process"][0]["save"]["hf_private"] = True
if concept_sentence:
config["config"]["process"][0]["trigger_word"] = concept_sentence
if sample_1 or sample_2 or sample_3:
config["config"]["process"][0]["train"]["disable_sampling"] = False
config["config"]["process"][0]["sample"]["sample_every"] = steps
config["config"]["process"][0]["sample"]["sample_steps"] = 28
config["config"]["process"][0]["sample"]["prompts"] = []
if sample_1:
config["config"]["process"][0]["sample"]["prompts"].append(sample_1)
if sample_2:
config["config"]["process"][0]["sample"]["prompts"].append(sample_2)
if sample_3:
config["config"]["process"][0]["sample"]["prompts"].append(sample_3)
else:
config["config"]["process"][0]["train"]["disable_sampling"] = True
if(model_to_train == "schnell"):
config["config"]["process"][0]["model"]["name_or_path"] = "black-forest-labs/FLUX.1-schnell"
config["config"]["process"][0]["model"]["assistant_lora_path"] = "ostris/FLUX.1-schnell-training-adapter"
config["config"]["process"][0]["sample"]["sample_steps"] = 4
if(use_more_advanced_options):
more_advanced_options_dict = yaml.safe_load(more_advanced_options)
config["config"]["process"][0] = recursive_update(config["config"]["process"][0], more_advanced_options_dict)
print(config)
# Save the updated config
# generate a random name for the config
random_config_name = str(uuid.uuid4())
os.makedirs("tmp", exist_ok=True)
config_path = f"tmp/{random_config_name}-{slugged_lora_name}.yaml"
with open(config_path, "w") as f:
yaml.dump(config, f)
# run the job locally
job = get_job(config_path)
job.run()
job.cleanup()
return f"Training completed successfully. Model saved as {slugged_lora_name}"
config_yaml = '''
device: cuda:0
model:
is_flux: true
quantize: true
network:
linear: 16 #it will overcome the 'rank' parameter
linear_alpha: 16 #you can have an alpha different than the ranking if you'd like
type: lora
sample:
guidance_scale: 3.5
height: 1024
neg: '' #doesn't work for FLUX
sample_every: 1000
sample_steps: 28
sampler: flowmatch
seed: 42
walk_seed: true
width: 1024
save:
dtype: float16
hf_private: true
max_step_saves_to_keep: 4
push_to_hub: true
save_every: 10000
train:
batch_size: 1
dtype: bf16
ema_config:
ema_decay: 0.99
use_ema: true
gradient_accumulation_steps: 1
gradient_checkpointing: true
noise_scheduler: flowmatch
optimizer: adamw8bit #options: prodigy, dadaptation, adamw, adamw8bit, lion, lion8bit
train_text_encoder: false #probably doesn't work for flux
train_unet: true
'''
theme = gr.themes.Monochrome(
text_size=gr.themes.Size(lg="18px", md="15px", sm="13px", xl="22px", xs="12px", xxl="24px", xxs="9px"),
font=[gr.themes.GoogleFont("Source Sans Pro"), "ui-sans-serif", "system-ui", "sans-serif"],
)
css = """
h1{font-size: 2em}
h3{margin-top: 0}
#component-1{text-align:center}
.main_ui_logged_out{opacity: 0.3; pointer-events: none}
.tabitem{border: 0px}
.group_padding{padding: .55em}
"""
with gr.Blocks(theme=theme, css=css) as demo:
gr.Markdown(
"""# LoRA Ease for FLUX 🧞‍♂️
### Train a high quality FLUX LoRA in a breeze ༄ using [Ostris' AI Toolkit](https://github.com/ostris/ai-toolkit)"""
)
with gr.Column() as main_ui:
with gr.Row():
lora_name = gr.Textbox(
label="The name of your LoRA",
info="This has to be a unique name",
placeholder="e.g.: Persian Miniature Painting style, Cat Toy",
)
concept_sentence = gr.Textbox(
label="Trigger word/sentence",
info="Trigger word or sentence to be used",
placeholder="uncommon word like p3rs0n or trtcrd, or sentence like 'in the style of CNSTLL'",
interactive=True,
)
with gr.Group(visible=True) as image_upload:
with gr.Row():
images = gr.File(
file_types=["image", ".txt"],
label="Upload your images",
file_count="multiple",
interactive=True,
visible=True,
scale=1,
)
with gr.Column(scale=3, visible=False) as captioning_area:
with gr.Column():
gr.Markdown(
"""# Custom captioning
<p style="margin-top:0">You can optionally add a custom caption for each image (or use an AI model for this). [trigger] will represent your concept sentence/trigger word.</p>
""", elem_classes="group_padding")
do_captioning = gr.Button("Add AI captions with Florence-2")
output_components = [captioning_area]
caption_list = []
for i in range(1, MAX_IMAGES + 1):
locals()[f"captioning_row_{i}"] = gr.Row(visible=False)
with locals()[f"captioning_row_{i}"]:
locals()[f"image_{i}"] = gr.Image(
type="filepath",
width=111,
height=111,
min_width=111,
interactive=False,
scale=2,
show_label=False,
show_share_button=False,
show_download_button=False,
)
locals()[f"caption_{i}"] = gr.Textbox(
label=f"Caption {i}", scale=15, interactive=True
)
output_components.append(locals()[f"captioning_row_{i}"])
output_components.append(locals()[f"image_{i}"])
output_components.append(locals()[f"caption_{i}"])
caption_list.append(locals()[f"caption_{i}"])
with gr.Accordion("Advanced options", open=False):
steps = gr.Number(label="Steps", value=1000, minimum=1, maximum=10000, step=1)
lr = gr.Number(label="Learning Rate", value=4e-4, minimum=1e-6, maximum=1e-3, step=1e-6)
rank = gr.Number(label="LoRA Rank", value=16, minimum=4, maximum=128, step=4)
model_to_train = gr.Radio(["dev", "schnell"], value="dev", label="Model to train")
low_vram = gr.Checkbox(label="Low VRAM", value=True)
with gr.Accordion("Even more advanced options", open=False):
use_more_advanced_options = gr.Checkbox(label="Use more advanced options", value=False)
more_advanced_options = gr.Code(config_yaml, language="yaml")
with gr.Accordion("Sample prompts (optional)", visible=False) as sample:
gr.Markdown(
"Include sample prompts to test out your trained model. Don't forget to include your trigger word/sentence (optional)"
)
sample_1 = gr.Textbox(label="Test prompt 1")
sample_2 = gr.Textbox(label="Test prompt 2")
sample_3 = gr.Textbox(label="Test prompt 3")
output_components.append(sample)
output_components.append(sample_1)
output_components.append(sample_2)
output_components.append(sample_3)
start = gr.Button("Start training", visible=False)
output_components.append(start)
progress_area = gr.Markdown("")
dataset_folder = gr.State()
images.upload(
load_captioning,
inputs=[images, concept_sentence],
outputs=output_components
)
images.delete(
load_captioning,
inputs=[images, concept_sentence],
outputs=output_components
)
images.clear(
hide_captioning,
outputs=[captioning_area, sample, start]
)
start.click(fn=create_dataset, inputs=[images] + caption_list, outputs=dataset_folder).then(
fn=start_training,
inputs=[
lora_name,
concept_sentence,
steps,
lr,
rank,
model_to_train,
low_vram,
dataset_folder,
sample_1,
sample_2,
sample_3,
use_more_advanced_options,
more_advanced_options
],
outputs=progress_area,
)
do_captioning.click(fn=run_captioning, inputs=[images, concept_sentence] + caption_list, outputs=caption_list)
if __name__ == "__main__":
demo.launch(share=True, show_error=True)

View File

@@ -24,6 +24,9 @@ class BaseProcess(object):
self.performance_log_every = self.get_conf('performance_log_every', 0)
print(json.dumps(self.config, indent=4))
def on_error(self, e: Exception):
pass
def get_conf(self, key, default=None, required=False, as_type=None):
# split key by '.' and recursively get the value

File diff suppressed because it is too large Load Diff

View File

@@ -1,7 +1,7 @@
import gc
import os
from collections import OrderedDict
from typing import ForwardRef, List
from typing import ForwardRef, List, Optional, Union
import torch
from safetensors.torch import save_file, load_file
@@ -22,6 +22,7 @@ class GenerateConfig:
self.sampler = kwargs.get('sampler', 'ddpm')
self.width = kwargs.get('width', 512)
self.height = kwargs.get('height', 512)
self.size_list: Union[List[int], None] = kwargs.get('size_list', None)
self.neg = kwargs.get('neg', '')
self.seed = kwargs.get('seed', -1)
self.guidance_scale = kwargs.get('guidance_scale', 7)
@@ -30,18 +31,34 @@ class GenerateConfig:
self.neg_2 = kwargs.get('neg_2', None)
self.prompts = kwargs.get('prompts', None)
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
self.compile = kwargs.get('compile', False)
self.ext = kwargs.get('ext', 'png')
self.prompt_file = kwargs.get('prompt_file', False)
self.num_repeats = kwargs.get('num_repeats', 1)
self.prompts_in_file = self.prompts
if self.prompts is None:
raise ValueError("Prompts must be set")
if isinstance(self.prompts, str):
if os.path.exists(self.prompts):
with open(self.prompts, 'r', encoding='utf-8') as f:
self.prompts = f.read().splitlines()
self.prompts = [p.strip() for p in self.prompts if len(p.strip()) > 0]
self.prompts_in_file = f.read().splitlines()
self.prompts_in_file = [p.strip() for p in self.prompts_in_file if len(p.strip()) > 0]
else:
raise ValueError("Prompts file does not exist, put in list if you want to use a list of prompts")
self.random_prompts = kwargs.get('random_prompts', False)
self.max_random_per_prompt = kwargs.get('max_random_per_prompt', 1)
self.max_images = kwargs.get('max_images', 10000)
if self.random_prompts:
self.prompts = []
for i in range(self.max_images):
num_prompts = random.randint(1, self.max_random_per_prompt)
prompt_list = [random.choice(self.prompts_in_file) for _ in range(num_prompts)]
self.prompts.append(", ".join(prompt_list))
else:
self.prompts = self.prompts_in_file
if kwargs.get('shuffle', False):
# shuffle the prompts
random.shuffle(self.prompts)
@@ -64,6 +81,7 @@ class GenerateProcess(BaseProcess):
self.model_config = ModelConfig(**self.get_conf('model', required=True))
self.device = self.get_conf('device', self.job.device)
self.generate_config = GenerateConfig(**self.get_conf('generate', required=True))
self.torch_dtype = get_torch_dtype(self.get_conf('dtype', 'float16'))
self.progress_bar = None
self.sd = StableDiffusion(
@@ -71,37 +89,58 @@ class GenerateProcess(BaseProcess):
model_config=self.model_config,
dtype=self.model_config.dtype,
)
print(f"Using device {self.device}")
def clean_prompt(self, prompt: str):
# remove any non alpha numeric characters or ,'" from prompt
return ''.join(e for e in prompt if e.isalnum() or e in ", '\"")
def run(self):
super().run()
print("Loading model...")
self.sd.load_model()
with torch.no_grad():
super().run()
print("Loading model...")
self.sd.load_model()
self.sd.pipeline.to(self.device, self.torch_dtype)
print(f"Generating {len(self.generate_config.prompts)} images")
# build prompt image configs
prompt_image_configs = []
for prompt in self.generate_config.prompts:
prompt_image_configs.append(GenerateImageConfig(
prompt=prompt,
prompt_2=self.generate_config.prompt_2,
width=self.generate_config.width,
height=self.generate_config.height,
num_inference_steps=self.generate_config.sample_steps,
guidance_scale=self.generate_config.guidance_scale,
negative_prompt=self.generate_config.neg,
negative_prompt_2=self.generate_config.neg_2,
seed=self.generate_config.seed,
guidance_rescale=self.generate_config.guidance_rescale,
output_ext=self.generate_config.ext,
output_folder=self.output_folder,
add_prompt_file=self.generate_config.prompt_file
))
# generate images
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
print("Compiling model...")
# self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead", fullgraph=True)
if self.generate_config.compile:
self.sd.unet = torch.compile(self.sd.unet, mode="reduce-overhead")
print("Done generating images")
# cleanup
del self.sd
gc.collect()
torch.cuda.empty_cache()
print(f"Generating {len(self.generate_config.prompts)} images")
# build prompt image configs
prompt_image_configs = []
for _ 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)
prompt_image_configs.append(GenerateImageConfig(
prompt=prompt,
prompt_2=self.generate_config.prompt_2,
width=width,
height=height,
num_inference_steps=self.generate_config.sample_steps,
guidance_scale=self.generate_config.guidance_scale,
negative_prompt=self.generate_config.neg,
negative_prompt_2=self.generate_config.neg_2,
seed=self.generate_config.seed,
guidance_rescale=self.generate_config.guidance_rescale,
output_ext=self.generate_config.ext,
output_folder=self.output_folder,
add_prompt_file=self.generate_config.prompt_file
))
# generate images
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
print("Done generating images")
# cleanup
del self.sd
gc.collect()
torch.cuda.empty_cache()

View File

@@ -275,6 +275,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
return adapter_tensors
def hook_train_loop(self, batch: Union['DataLoaderBatchDTO', None]):
if isinstance(batch, list):
batch = batch[0]
# set to eval mode
self.sd.set_device_state(self.eval_slider_device_state)
with torch.no_grad():
@@ -364,14 +366,36 @@ class TrainSliderProcess(BaseSDTrainProcess):
denoised_latents = torch.cat([noisy_latents] * self.prompt_chunk_size, dim=0)
current_timestep = timesteps
else:
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
if self.train_config.noise_scheduler == 'flowmatch':
linear_timesteps = any([
self.train_config.linear_timesteps,
self.train_config.linear_timesteps2,
self.train_config.timestep_type == 'linear',
])
timestep_type = 'linear' if linear_timesteps else None
if timestep_type is None:
timestep_type = self.train_config.timestep_type
# make fake latents
l = torch.randn(
true_batch_size, 16, height, width
).to(self.device_torch, dtype=dtype)
self.sd.noise_scheduler.set_train_timesteps(
self.train_config.max_denoising_steps,
device=self.device_torch,
timestep_type=timestep_type,
latents=l
)
else:
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
# ger a random number of steps
timesteps_to = torch.randint(
1, self.train_config.max_denoising_steps, (1,)
1, self.train_config.max_denoising_steps - 1, (1,)
).item()
# get noise
@@ -389,7 +413,8 @@ class TrainSliderProcess(BaseSDTrainProcess):
assert not self.network.is_active
self.sd.unet.eval()
# pass the multiplier list to the network
self.network.multiplier = prompt_pair.multiplier_list
# double up since we are doing cfg
self.network.multiplier = prompt_pair.multiplier_list + prompt_pair.multiplier_list
denoised_latents = self.sd.diffuse_some_steps(
latents, # pass simple noise latents
train_tools.concat_prompt_embeddings(
@@ -507,7 +532,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
for anchor_chunk, denoised_latent_chunk, anchor_target_noise_chunk in zip(
anchor_chunks, denoised_latent_chunks, anchor_target_noise_chunks
):
self.network.multiplier = anchor_chunk.multiplier_list
self.network.multiplier = anchor_chunk.multiplier_list + anchor_chunk.multiplier_list
anchor_pred_noise = get_noise_pred(
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
@@ -582,7 +607,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
mask_multiplier_chunks,
unmasked_target_chunks
):
self.network.multiplier = prompt_pair_chunk.multiplier_list
self.network.multiplier = prompt_pair_chunk.multiplier_list + prompt_pair_chunk.multiplier_list
target_latents = get_noise_pred(
prompt_pair_chunk.positive_target,
prompt_pair_chunk.target_class,
@@ -611,6 +636,7 @@ class TrainSliderProcess(BaseSDTrainProcess):
offset_neutral = neutral_latents_chunk
# offsets are already adjusted on a per-batch basis
offset_neutral += offset
offset_neutral = offset_neutral.detach().requires_grad_(False)
# 16.15 GB RAM for 512x512 -> 4.20GB RAM for 512x512 with new grad_checkpointing
loss = torch.nn.functional.mse_loss(target_latents.float(), offset_neutral.float(), reduction="none")

View File

@@ -1,6 +1,7 @@
import copy
import glob
import os
import shutil
import time
from collections import OrderedDict
@@ -13,6 +14,7 @@ from torch import nn
from torchvision.transforms import transforms
from jobs.process import BaseTrainProcess
from toolkit.image_utils import show_tensors
from toolkit.kohya_model_util import load_vae, convert_diffusers_back_to_ldm
from toolkit.data_loader import ImageDataset
from toolkit.losses import ComparativeTotalVariation, get_gradient_penalty, PatternLoss
@@ -25,6 +27,8 @@ from tqdm import tqdm
import time
import numpy as np
from .models.vgg19_critic import Critic
from torchvision.transforms import Resize
import lpips
IMAGE_TRANSFORMS = transforms.Compose(
[
@@ -62,6 +66,7 @@ class TrainVAEProcess(BaseTrainProcess):
self.kld_weight = self.get_conf('kld_weight', 0, as_type=float)
self.mse_weight = self.get_conf('mse_weight', 1e0, as_type=float)
self.tv_weight = self.get_conf('tv_weight', 1e0, as_type=float)
self.lpips_weight = self.get_conf('lpips_weight', 1e0, as_type=float)
self.critic_weight = self.get_conf('critic_weight', 1, as_type=float)
self.pattern_weight = self.get_conf('pattern_weight', 1, as_type=float)
self.optimizer_params = self.get_conf('optimizer_params', {})
@@ -71,6 +76,9 @@ class TrainVAEProcess(BaseTrainProcess):
self.vgg_19 = None
self.style_weight_scalers = []
self.content_weight_scalers = []
self.lpips_loss:lpips.LPIPS = None
self.vae_scale_factor = 8
self.step_num = 0
self.epoch_num = 0
@@ -137,6 +145,15 @@ class TrainVAEProcess(BaseTrainProcess):
num_workers=6
)
def remove_oldest_checkpoint(self):
max_to_keep = 4
folders = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
if len(folders) > max_to_keep:
folders.sort(key=os.path.getmtime)
for folder in folders[:-max_to_keep]:
print(f"Removing {folder}")
shutil.rmtree(folder)
def setup_vgg19(self):
if self.vgg_19 is None:
self.vgg_19, self.style_losses, self.content_losses, self.vgg19_pool_4 = get_style_model_and_losses(
@@ -211,7 +228,7 @@ class TrainVAEProcess(BaseTrainProcess):
def get_pattern_loss(self, pred, target):
if self._pattern_loss is None:
self._pattern_loss = PatternLoss(pattern_size=8, dtype=self.torch_dtype).to(self.device,
self._pattern_loss = PatternLoss(pattern_size=16, dtype=self.torch_dtype).to(self.device,
dtype=self.torch_dtype)
loss = torch.mean(self._pattern_loss(pred, target))
return loss
@@ -226,25 +243,21 @@ class TrainVAEProcess(BaseTrainProcess):
step_num = f"_{str(step).zfill(9)}"
self.update_training_metadata()
filename = f'{self.job.name}{step_num}.safetensors'
# prepare meta
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
filename = f'{self.job.name}{step_num}_diffusers'
state_dict = convert_diffusers_back_to_ldm(self.vae)
for key in list(state_dict.keys()):
v = state_dict[key]
v = v.detach().clone().to("cpu").to(torch.float32)
state_dict[key] = v
# having issues with meta
save_file(state_dict, os.path.join(self.save_root, filename), save_meta)
self.vae = self.vae.to("cpu", dtype=torch.float16)
self.vae.save_pretrained(
save_directory=os.path.join(self.save_root, filename)
)
self.vae = self.vae.to(self.device, dtype=self.torch_dtype)
self.print(f"Saved to {os.path.join(self.save_root, filename)}")
if self.use_critic:
self.critic.save(step)
self.remove_oldest_checkpoint()
def sample(self, step=None):
sample_folder = os.path.join(self.save_root, 'samples')
if not os.path.exists(sample_folder):
@@ -280,6 +293,13 @@ class TrainVAEProcess(BaseTrainProcess):
output_img.paste(input_img, (0, 0))
output_img.paste(decoded, (self.resolution, 0))
scale_up = 2
if output_img.height <= 300:
scale_up = 4
# scale up using nearest neighbor
output_img = output_img.resize((output_img.width * scale_up, output_img.height * scale_up), Image.NEAREST)
step_num = ''
if step is not None:
# zero-pad 9 digits
@@ -294,7 +314,7 @@ class TrainVAEProcess(BaseTrainProcess):
path_to_load = self.vae_path
# see if we have a checkpoint in out output to resume from
self.print(f"Looking for latest checkpoint in {self.save_root}")
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*.safetensors"))
files = glob.glob(os.path.join(self.save_root, f"{self.job.name}*_diffusers"))
if files and len(files) > 0:
latest_file = max(files, key=os.path.getmtime)
print(f" - Latest checkpoint is: {latest_file}")
@@ -306,13 +326,14 @@ class TrainVAEProcess(BaseTrainProcess):
self.print(f"Loading VAE")
self.print(f" - Loading VAE: {path_to_load}")
if self.vae is None:
self.vae = load_vae(path_to_load, dtype=self.torch_dtype)
self.vae = AutoencoderKL.from_pretrained(path_to_load)
# set decoder to train
self.vae.to(self.device, dtype=self.torch_dtype)
self.vae.requires_grad_(False)
self.vae.eval()
self.vae.decoder.train()
self.vae_scale_factor = 2 ** (len(self.vae.config['block_out_channels']) - 1)
def run(self):
super().run()
@@ -374,6 +395,10 @@ class TrainVAEProcess(BaseTrainProcess):
if self.use_critic:
self.critic.setup()
if self.lpips_weight > 0 and self.lpips_loss is None:
# self.lpips_loss = lpips.LPIPS(net='vgg')
self.lpips_loss = lpips.LPIPS(net='vgg').to(self.device, dtype=self.torch_dtype)
optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
optimizer_params=self.optimizer_params)
@@ -397,6 +422,7 @@ class TrainVAEProcess(BaseTrainProcess):
self.sample()
blank_losses = OrderedDict({
"total": [],
"lpips": [],
"style": [],
"content": [],
"mse": [],
@@ -415,17 +441,29 @@ class TrainVAEProcess(BaseTrainProcess):
for batch in self.data_loader:
if self.step_num >= self.max_steps:
break
with torch.no_grad():
batch = batch.to(self.device, dtype=self.torch_dtype)
batch = batch.to(self.device, dtype=self.torch_dtype)
# forward pass
dgd = self.vae.encode(batch).latent_dist
mu, logvar = dgd.mean, dgd.logvar
latents = dgd.sample()
latents.requires_grad_(True)
# resize so it matches size of vae evenly
if batch.shape[2] % self.vae_scale_factor != 0 or batch.shape[3] % self.vae_scale_factor != 0:
batch = Resize((batch.shape[2] // self.vae_scale_factor * self.vae_scale_factor,
batch.shape[3] // self.vae_scale_factor * self.vae_scale_factor))(batch)
# forward pass
dgd = self.vae.encode(batch).latent_dist
mu, logvar = dgd.mean, dgd.logvar
latents = dgd.sample()
latents.detach().requires_grad_(True)
pred = self.vae.decode(latents).sample
with torch.no_grad():
show_tensors(
pred.clamp(-1, 1).clone(),
"combined tensor"
)
# Run through VGG19
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
stacked = torch.cat([pred, batch], dim=0)
@@ -441,14 +479,31 @@ class TrainVAEProcess(BaseTrainProcess):
content_loss = self.get_content_loss() * self.content_weight
kld_loss = self.get_kld_loss(mu, logvar) * self.kld_weight
mse_loss = self.get_mse_loss(pred, batch) * self.mse_weight
if self.lpips_weight > 0:
lpips_loss = self.lpips_loss(
pred.clamp(-1, 1),
batch.clamp(-1, 1)
).mean() * self.lpips_weight
else:
lpips_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
tv_loss = self.get_tv_loss(pred, batch) * self.tv_weight
pattern_loss = self.get_pattern_loss(pred, batch) * self.pattern_weight
if self.use_critic:
critic_gen_loss = self.critic.get_critic_loss(self.vgg19_pool_4.tensor) * self.critic_weight
# do not let abs critic gen loss be higher than abs lpips * 0.1 if using it
if self.lpips_weight > 0:
max_target = lpips_loss.abs() * 0.1
with torch.no_grad():
crit_g_scaler = 1.0
if critic_gen_loss.abs() > max_target:
crit_g_scaler = max_target / critic_gen_loss.abs()
critic_gen_loss *= crit_g_scaler
else:
critic_gen_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss
loss = style_loss + content_loss + kld_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss + lpips_loss
# Backward pass and optimization
optimizer.zero_grad()
@@ -460,6 +515,8 @@ class TrainVAEProcess(BaseTrainProcess):
loss_value = loss.item()
# get exponent like 3.54e-4
loss_string = f"loss: {loss_value:.2e}"
if self.lpips_weight > 0:
loss_string += f" lpips: {lpips_loss.item():.2e}"
if self.content_weight > 0:
loss_string += f" cnt: {content_loss.item():.2e}"
if self.style_weight > 0:
@@ -477,7 +534,8 @@ class TrainVAEProcess(BaseTrainProcess):
if self.use_critic:
loss_string += f" crD: {critic_d_loss:.2e}"
if self.optimizer_type.startswith('dadaptation'):
if self.optimizer_type.startswith('dadaptation') or \
self.optimizer_type.lower().startswith('prodigy'):
learning_rate = (
optimizer.param_groups[0]["d"] *
optimizer.param_groups[0]["lr"]
@@ -495,6 +553,7 @@ class TrainVAEProcess(BaseTrainProcess):
self.progress_bar.update(1)
epoch_losses["total"].append(loss_value)
epoch_losses["lpips"].append(lpips_loss.item())
epoch_losses["style"].append(style_loss.item())
epoch_losses["content"].append(content_loss.item())
epoch_losses["mse"].append(mse_loss.item())
@@ -505,6 +564,7 @@ class TrainVAEProcess(BaseTrainProcess):
epoch_losses["crD"].append(critic_d_loss)
log_losses["total"].append(loss_value)
log_losses["lpips"].append(lpips_loss.item())
log_losses["style"].append(style_loss.item())
log_losses["content"].append(content_loss.item())
log_losses["mse"].append(mse_loss.item())

View File

@@ -0,0 +1,291 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"collapsed": false,
"id": "zl-S0m3pkQC5"
},
"source": [
"# AI Toolkit by Ostris\n",
"## FLUX.1-dev Training\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"!nvidia-smi"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "BvAG0GKAh59G"
},
"outputs": [],
"source": [
"!git clone https://github.com/ostris/ai-toolkit\n",
"!mkdir -p /content/dataset"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "UFUW4ZMmnp1V"
},
"source": [
"Put your image dataset in the `/content/dataset` folder"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "XGZqVER_aQJW"
},
"outputs": [],
"source": [
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "OV0HnOI6o8V6"
},
"source": [
"## Model License\n",
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
"\n",
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
"\n",
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "3yZZdhFRoj2m"
},
"outputs": [],
"source": [
"import getpass\n",
"import os\n",
"\n",
"# Prompt for the token\n",
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
"\n",
"# Set the environment variable\n",
"os.environ['HF_TOKEN'] = hf_token\n",
"\n",
"print(\"HF_TOKEN environment variable has been set.\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "9gO2EzQ1kQC8"
},
"outputs": [],
"source": [
"import os\n",
"import sys\n",
"sys.path.append('/content/ai-toolkit')\n",
"from toolkit.job import run_job\n",
"from collections import OrderedDict\n",
"from PIL import Image\n",
"import os\n",
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "N8UUFzVRigbC"
},
"source": [
"## Setup\n",
"\n",
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "_t28QURYjRQO"
},
"outputs": [],
"source": [
"from collections import OrderedDict\n",
"\n",
"job_to_run = OrderedDict([\n",
" ('job', 'extension'),\n",
" ('config', OrderedDict([\n",
" # this name will be the folder and filename name\n",
" ('name', 'my_first_flux_lora_v1'),\n",
" ('process', [\n",
" OrderedDict([\n",
" ('type', 'sd_trainer'),\n",
" # root folder to save training sessions/samples/weights\n",
" ('training_folder', '/content/output'),\n",
" # uncomment to see performance stats in the terminal every N steps\n",
" #('performance_log_every', 1000),\n",
" ('device', 'cuda:0'),\n",
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
" # ('trigger_word', 'image'),\n",
" ('network', OrderedDict([\n",
" ('type', 'lora'),\n",
" ('linear', 16),\n",
" ('linear_alpha', 16)\n",
" ])),\n",
" ('save', OrderedDict([\n",
" ('dtype', 'float16'), # precision to save\n",
" ('save_every', 250), # save every this many steps\n",
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
" ])),\n",
" ('datasets', [\n",
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
" # images will automatically be resized and bucketed into the resolution specified\n",
" OrderedDict([\n",
" ('folder_path', '/content/dataset'),\n",
" ('caption_ext', 'txt'),\n",
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
" ])\n",
" ]),\n",
" ('train', OrderedDict([\n",
" ('batch_size', 1),\n",
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
" ('gradient_accumulation_steps', 1),\n",
" ('train_unet', True),\n",
" ('train_text_encoder', False), # probably won't work with flux\n",
" ('content_or_style', 'balanced'), # content, style, balanced\n",
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
" ('noise_scheduler', 'flowmatch'), # for training only\n",
" ('optimizer', 'adamw8bit'),\n",
" ('lr', 1e-4),\n",
"\n",
" # uncomment this to skip the pre training sample\n",
" # ('skip_first_sample', True),\n",
"\n",
" # uncomment to completely disable sampling\n",
" # ('disable_sampling', True),\n",
"\n",
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
" # ('linear_timesteps', True),\n",
"\n",
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
" ('ema_config', OrderedDict([\n",
" ('use_ema', True),\n",
" ('ema_decay', 0.99)\n",
" ])),\n",
"\n",
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
" ('dtype', 'bf16')\n",
" ])),\n",
" ('model', OrderedDict([\n",
" # huggingface model name or path\n",
" ('name_or_path', 'black-forest-labs/FLUX.1-dev'),\n",
" ('is_flux', True),\n",
" ('quantize', True), # run 8bit mixed precision\n",
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
" ])),\n",
" ('sample', OrderedDict([\n",
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
" ('sample_every', 250), # sample every this many steps\n",
" ('width', 1024),\n",
" ('height', 1024),\n",
" ('prompts', [\n",
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
" 'a bear building a log cabin in the snow covered mountains',\n",
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
" 'hipster man with a beard, building a chair, in a wood shop',\n",
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
" ]),\n",
" ('neg', ''), # not used on flux\n",
" ('seed', 42),\n",
" ('walk_seed', True),\n",
" ('guidance_scale', 4),\n",
" ('sample_steps', 20)\n",
" ]))\n",
" ])\n",
" ])\n",
" ])),\n",
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
" ('meta', OrderedDict([\n",
" ('name', '[name]'),\n",
" ('version', '1.0')\n",
" ]))\n",
"])\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "h6F1FlM2Wb3l"
},
"source": [
"## Run it\n",
"\n",
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "HkajwI8gteOh"
},
"outputs": [],
"source": [
"run_job(job_to_run)\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "Hblgb5uwW5SD"
},
"source": [
"## Done\n",
"\n",
"Check your ourput dir and get your slider\n"
]
}
],
"metadata": {
"accelerator": "GPU",
"colab": {
"gpuType": "A100",
"machine_shape": "hm",
"provenance": []
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
},
"language_info": {
"name": "python"
}
},
"nbformat": 4,
"nbformat_minor": 0
}

View File

@@ -0,0 +1,296 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"collapsed": false,
"id": "zl-S0m3pkQC5"
},
"source": [
"# AI Toolkit by Ostris\n",
"## FLUX.1-schnell Training\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "3cokMT-WC6rG"
},
"outputs": [],
"source": [
"!nvidia-smi"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"collapsed": true,
"id": "BvAG0GKAh59G"
},
"outputs": [],
"source": [
"!git clone https://github.com/ostris/ai-toolkit\n",
"!mkdir -p /content/dataset"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "UFUW4ZMmnp1V"
},
"source": [
"Put your image dataset in the `/content/dataset` folder"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"collapsed": true,
"id": "XGZqVER_aQJW"
},
"outputs": [],
"source": [
"!cd ai-toolkit && git submodule update --init --recursive && pip install -r requirements.txt\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "OV0HnOI6o8V6"
},
"source": [
"## Model License\n",
"Training currently only works with FLUX.1-dev. Which means anything you train will inherit the non-commercial license. It is also a gated model, so you need to accept the license on HF before using it. Otherwise, this will fail. Here are the required steps to setup a license.\n",
"\n",
"Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)\n",
"\n",
"[Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and place it in the next cell after running it."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "3yZZdhFRoj2m"
},
"outputs": [],
"source": [
"import getpass\n",
"import os\n",
"\n",
"# Prompt for the token\n",
"hf_token = getpass.getpass('Enter your HF access token and press enter: ')\n",
"\n",
"# Set the environment variable\n",
"os.environ['HF_TOKEN'] = hf_token\n",
"\n",
"print(\"HF_TOKEN environment variable has been set.\")"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {
"id": "9gO2EzQ1kQC8"
},
"outputs": [],
"source": [
"import os\n",
"import sys\n",
"sys.path.append('/content/ai-toolkit')\n",
"from toolkit.job import run_job\n",
"from collections import OrderedDict\n",
"from PIL import Image\n",
"import os\n",
"os.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"1\""
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "N8UUFzVRigbC"
},
"source": [
"## Setup\n",
"\n",
"This is your config. It is documented pretty well. Normally you would do this as a yaml file, but for colab, this will work. This will run as is without modification, but feel free to edit as you want."
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {
"id": "_t28QURYjRQO"
},
"outputs": [],
"source": [
"from collections import OrderedDict\n",
"\n",
"job_to_run = OrderedDict([\n",
" ('job', 'extension'),\n",
" ('config', OrderedDict([\n",
" # this name will be the folder and filename name\n",
" ('name', 'my_first_flux_lora_v1'),\n",
" ('process', [\n",
" OrderedDict([\n",
" ('type', 'sd_trainer'),\n",
" # root folder to save training sessions/samples/weights\n",
" ('training_folder', '/content/output'),\n",
" # uncomment to see performance stats in the terminal every N steps\n",
" #('performance_log_every', 1000),\n",
" ('device', 'cuda:0'),\n",
" # if a trigger word is specified, it will be added to captions of training data if it does not already exist\n",
" # alternatively, in your captions you can add [trigger] and it will be replaced with the trigger word\n",
" # ('trigger_word', 'image'),\n",
" ('network', OrderedDict([\n",
" ('type', 'lora'),\n",
" ('linear', 16),\n",
" ('linear_alpha', 16)\n",
" ])),\n",
" ('save', OrderedDict([\n",
" ('dtype', 'float16'), # precision to save\n",
" ('save_every', 250), # save every this many steps\n",
" ('max_step_saves_to_keep', 4) # how many intermittent saves to keep\n",
" ])),\n",
" ('datasets', [\n",
" # datasets are a folder of images. captions need to be txt files with the same name as the image\n",
" # for instance image2.jpg and image2.txt. Only jpg, jpeg, and png are supported currently\n",
" # images will automatically be resized and bucketed into the resolution specified\n",
" OrderedDict([\n",
" ('folder_path', '/content/dataset'),\n",
" ('caption_ext', 'txt'),\n",
" ('caption_dropout_rate', 0.05), # will drop out the caption 5% of time\n",
" ('shuffle_tokens', False), # shuffle caption order, split by commas\n",
" ('cache_latents_to_disk', True), # leave this true unless you know what you're doing\n",
" ('resolution', [512, 768, 1024]) # flux enjoys multiple resolutions\n",
" ])\n",
" ]),\n",
" ('train', OrderedDict([\n",
" ('batch_size', 1),\n",
" ('steps', 2000), # total number of steps to train 500 - 4000 is a good range\n",
" ('gradient_accumulation_steps', 1),\n",
" ('train_unet', True),\n",
" ('train_text_encoder', False), # probably won't work with flux\n",
" ('gradient_checkpointing', True), # need the on unless you have a ton of vram\n",
" ('noise_scheduler', 'flowmatch'), # for training only\n",
" ('optimizer', 'adamw8bit'),\n",
" ('lr', 1e-4),\n",
"\n",
" # uncomment this to skip the pre training sample\n",
" # ('skip_first_sample', True),\n",
"\n",
" # uncomment to completely disable sampling\n",
" # ('disable_sampling', True),\n",
"\n",
" # uncomment to use new vell curved weighting. Experimental but may produce better results\n",
" # ('linear_timesteps', True),\n",
"\n",
" # ema will smooth out learning, but could slow it down. Recommended to leave on.\n",
" ('ema_config', OrderedDict([\n",
" ('use_ema', True),\n",
" ('ema_decay', 0.99)\n",
" ])),\n",
"\n",
" # will probably need this if gpu supports it for flux, other dtypes may not work correctly\n",
" ('dtype', 'bf16')\n",
" ])),\n",
" ('model', OrderedDict([\n",
" # huggingface model name or path\n",
" ('name_or_path', 'black-forest-labs/FLUX.1-schnell'),\n",
" ('assistant_lora_path', 'ostris/FLUX.1-schnell-training-adapter'), # Required for flux schnell training\n",
" ('is_flux', True),\n",
" ('quantize', True), # run 8bit mixed precision\n",
" # low_vram is painfully slow to fuse in the adapter avoid it unless absolutely necessary\n",
" #('low_vram', True), # uncomment this if the GPU is connected to your monitors. It will use less vram to quantize, but is slower.\n",
" ])),\n",
" ('sample', OrderedDict([\n",
" ('sampler', 'flowmatch'), # must match train.noise_scheduler\n",
" ('sample_every', 250), # sample every this many steps\n",
" ('width', 1024),\n",
" ('height', 1024),\n",
" ('prompts', [\n",
" # you can add [trigger] to the prompts here and it will be replaced with the trigger word\n",
" #'[trigger] holding a sign that says \\'I LOVE PROMPTS!\\'',\n",
" 'woman with red hair, playing chess at the park, bomb going off in the background',\n",
" 'a woman holding a coffee cup, in a beanie, sitting at a cafe',\n",
" 'a horse is a DJ at a night club, fish eye lens, smoke machine, lazer lights, holding a martini',\n",
" 'a man showing off his cool new t shirt at the beach, a shark is jumping out of the water in the background',\n",
" 'a bear building a log cabin in the snow covered mountains',\n",
" 'woman playing the guitar, on stage, singing a song, laser lights, punk rocker',\n",
" 'hipster man with a beard, building a chair, in a wood shop',\n",
" 'photo of a man, white background, medium shot, modeling clothing, studio lighting, white backdrop',\n",
" 'a man holding a sign that says, \\'this is a sign\\'',\n",
" 'a bulldog, in a post apocalyptic world, with a shotgun, in a leather jacket, in a desert, with a motorcycle'\n",
" ]),\n",
" ('neg', ''), # not used on flux\n",
" ('seed', 42),\n",
" ('walk_seed', True),\n",
" ('guidance_scale', 1), # schnell does not do guidance\n",
" ('sample_steps', 4) # 1 - 4 works well\n",
" ]))\n",
" ])\n",
" ])\n",
" ])),\n",
" # you can add any additional meta info here. [name] is replaced with config name at top\n",
" ('meta', OrderedDict([\n",
" ('name', '[name]'),\n",
" ('version', '1.0')\n",
" ]))\n",
"])\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "h6F1FlM2Wb3l"
},
"source": [
"## Run it\n",
"\n",
"Below does all the magic. Check your folders to the left. Items will be in output/LoRA/your_name_v1 In the samples folder, there are preiodic sampled. This doesnt work great with colab. They will be in /content/output"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {
"id": "HkajwI8gteOh"
},
"outputs": [],
"source": [
"run_job(job_to_run)\n"
]
},
{
"cell_type": "markdown",
"metadata": {
"id": "Hblgb5uwW5SD"
},
"source": [
"## Done\n",
"\n",
"Check your ourput dir and get your slider\n"
]
}
],
"metadata": {
"accelerator": "GPU",
"colab": {
"gpuType": "A100",
"machine_shape": "hm",
"provenance": []
},
"kernelspec": {
"display_name": "Python 3",
"name": "python3"
},
"language_info": {
"name": "python"
}
},
"nbformat": 4,
"nbformat_minor": 0
}

View File

@@ -1,8 +1,9 @@
torch
torchvision
torch==2.6.0
torchvision==0.21.0
torchao==0.9.0
safetensors
diffusers==0.21.3
git+https://github.com/huggingface/transformers.git
git+https://github.com/huggingface/diffusers@363d1ab7e24c5ed6c190abb00df66d9edb74383b
transformers==4.49.0
lycoris-lora==1.8.3
flatten_json
pyyaml
@@ -13,7 +14,8 @@ invisible-watermark
einops
accelerate
toml
albumentations
albumentations==1.4.15
albucore==0.0.16
pydantic
omegaconf
k-diffusion
@@ -21,4 +23,16 @@ open_clip_torch
timm
prodigyopt
controlnet_aux==0.0.7
python-dotenv
python-dotenv
bitsandbytes
hf_transfer
lpips
pytorch_fid
optimum-quanto==0.2.4
sentencepiece
huggingface_hub
peft
gradio
python-slugify
opencv-python
pytorch-wavelets==1.3.0

46
run.py
View File

@@ -1,4 +1,5 @@
import os
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
import sys
from typing import Union, OrderedDict
from dotenv import load_dotenv
@@ -19,20 +20,26 @@ if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
torch.autograd.set_detect_anomaly(True)
import argparse
from toolkit.job import get_job
from toolkit.accelerator import get_accelerator
from toolkit.print import print_acc, setup_log_to_file
accelerator = get_accelerator()
def print_end_message(jobs_completed, jobs_failed):
if not accelerator.is_main_process:
return
failure_string = f"{jobs_failed} failure{'' if jobs_failed == 1 else 's'}" if jobs_failed > 0 else ""
completed_string = f"{jobs_completed} completed job{'' if jobs_completed == 1 else 's'}"
print("")
print("========================================")
print("Result:")
print_acc("")
print_acc("========================================")
print_acc("Result:")
if len(completed_string) > 0:
print(f" - {completed_string}")
print_acc(f" - {completed_string}")
if len(failure_string) > 0:
print(f" - {failure_string}")
print("========================================")
print_acc(f" - {failure_string}")
print_acc("========================================")
def main():
@@ -60,7 +67,17 @@ def main():
default=None,
help='Name to replace [name] tag in config file, useful for shared config file'
)
parser.add_argument(
'-l', '--log',
type=str,
default=None,
help='Log file to write output to'
)
args = parser.parse_args()
if args.log is not None:
setup_log_to_file(args.log)
config_file_list = args.config_file_list
if len(config_file_list) == 0:
@@ -69,7 +86,8 @@ def main():
jobs_completed = 0
jobs_failed = 0
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
if accelerator.is_main_process:
print_acc(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
for config_file in config_file_list:
try:
@@ -78,8 +96,20 @@ def main():
job.cleanup()
jobs_completed += 1
except Exception as e:
print(f"Error running job: {e}")
print_acc(f"Error running job: {e}")
jobs_failed += 1
try:
job.process[0].on_error(e)
except Exception as e2:
print_acc(f"Error running on_error: {e2}")
if not args.recover:
print_end_message(jobs_completed, jobs_failed)
raise e
except KeyboardInterrupt as e:
try:
job.process[0].on_error(e)
except Exception as e2:
print_acc(f"Error running on_error: {e2}")
if not args.recover:
print_end_message(jobs_completed, jobs_failed)
raise e

175
run_modal.py Normal file
View File

@@ -0,0 +1,175 @@
'''
ostris/ai-toolkit on https://modal.com
Run training with the following command:
modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml
'''
import os
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
import sys
import modal
from dotenv import load_dotenv
# Load the .env file if it exists
load_dotenv()
sys.path.insert(0, "/root/ai-toolkit")
# must come before ANY torch or fastai imports
# import toolkit.cuda_malloc
# turn off diffusers telemetry until I can figure out how to make it opt-in
os.environ['DISABLE_TELEMETRY'] = 'YES'
# define the volume for storing model outputs, using "creating volumes lazily": https://modal.com/docs/guide/volumes
# you will find your model, samples and optimizer stored in: https://modal.com/storage/your-username/main/flux-lora-models
model_volume = modal.Volume.from_name("flux-lora-models", create_if_missing=True)
# modal_output, due to "cannot mount volume on non-empty path" requirement
MOUNT_DIR = "/root/ai-toolkit/modal_output" # modal_output, due to "cannot mount volume on non-empty path" requirement
# define modal app
image = (
modal.Image.debian_slim(python_version="3.11")
# install required system and pip packages, more about this modal approach: https://modal.com/docs/examples/dreambooth_app
.apt_install("libgl1", "libglib2.0-0")
.pip_install(
"python-dotenv",
"torch",
"diffusers[torch]",
"transformers",
"ftfy",
"torchvision",
"oyaml",
"opencv-python",
"albumentations",
"safetensors",
"lycoris-lora==1.8.3",
"flatten_json",
"pyyaml",
"tensorboard",
"kornia",
"invisible-watermark",
"einops",
"accelerate",
"toml",
"pydantic",
"omegaconf",
"k-diffusion",
"open_clip_torch",
"timm",
"prodigyopt",
"controlnet_aux==0.0.7",
"bitsandbytes",
"hf_transfer",
"lpips",
"pytorch_fid",
"optimum-quanto",
"sentencepiece",
"huggingface_hub",
"peft"
)
)
# mount for the entire ai-toolkit directory
# example: "/Users/username/ai-toolkit" is the local directory, "/root/ai-toolkit" is the remote directory
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
# create the Modal app with the necessary mounts and volumes
app = modal.App(name="flux-lora-training", image=image, mounts=[code_mount], volumes={MOUNT_DIR: model_volume})
# Check if we have DEBUG_TOOLKIT in env
if os.environ.get("DEBUG_TOOLKIT", "0") == "1":
# Set torch to trace mode
import torch
torch.autograd.set_detect_anomaly(True)
import argparse
from toolkit.job import get_job
def print_end_message(jobs_completed, jobs_failed):
failure_string = f"{jobs_failed} failure{'' if jobs_failed == 1 else 's'}" if jobs_failed > 0 else ""
completed_string = f"{jobs_completed} completed job{'' if jobs_completed == 1 else 's'}"
print("")
print("========================================")
print("Result:")
if len(completed_string) > 0:
print(f" - {completed_string}")
if len(failure_string) > 0:
print(f" - {failure_string}")
print("========================================")
@app.function(
# request a GPU with at least 24GB VRAM
# more about modal GPU's: https://modal.com/docs/guide/gpu
gpu="A100", # gpu="H100"
# more about modal timeouts: https://modal.com/docs/guide/timeouts
timeout=7200 # 2 hours, increase or decrease if needed
)
def main(config_file_list_str: str, recover: bool = False, name: str = None):
# convert the config file list from a string to a list
config_file_list = config_file_list_str.split(",")
jobs_completed = 0
jobs_failed = 0
print(f"Running {len(config_file_list)} job{'' if len(config_file_list) == 1 else 's'}")
for config_file in config_file_list:
try:
job = get_job(config_file, name)
job.config['process'][0]['training_folder'] = MOUNT_DIR
os.makedirs(MOUNT_DIR, exist_ok=True)
print(f"Training outputs will be saved to: {MOUNT_DIR}")
# run the job
job.run()
# commit the volume after training
model_volume.commit()
job.cleanup()
jobs_completed += 1
except Exception as e:
print(f"Error running job: {e}")
jobs_failed += 1
if not recover:
print_end_message(jobs_completed, jobs_failed)
raise e
print_end_message(jobs_completed, jobs_failed)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# require at least one config file
parser.add_argument(
'config_file_list',
nargs='+',
type=str,
help='Name of config file (eg: person_v1 for config/person_v1.json/yaml), or full path if it is not in config folder, you can pass multiple config files and run them all sequentially'
)
# flag to continue if a job fails
parser.add_argument(
'-r', '--recover',
action='store_true',
help='Continue running additional jobs even if a job fails'
)
# optional name replacement for config file
parser.add_argument(
'-n', '--name',
type=str,
default=None,
help='Name to replace [name] tag in config file, useful for shared config file'
)
args = parser.parse_args()
# convert list of config files to a comma-separated string for Modal compatibility
config_file_list_str = ",".join(args.config_file_list)
main.call(config_file_list_str=config_file_list_str, recover=args.recover, name=args.name)

View File

@@ -0,0 +1,426 @@
#######################################################
# Convert Diffusers Flux/Flex to all in one ComfyUI safetensors file
# The VAE, T5 and clip will all be in the safetensors file
# T5 will always be 8bit with the all in one file
# You can save the transformer weights as bf16 or 8-bit with the --do_8_bit flag
#
# Download a reference model from Huggingface
# https://huggingface.co/Comfy-Org/flux1-dev/blob/main/flux1-dev-fp8.safetensors
#
# Call like this for 8-bit transformer weights:
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors --do_8_bit
#
# Call like this for bf16 transformer weights:
# python convert_flux_diffusers_to_orig.py /path/to/diffusers/checkpoint /path/to/flux1-dev-fp8.safetensors /output/path/my_finetune.safetensors
#
#######################################################
import argparse
from datetime import date
import json
import os
from pathlib import Path
import safetensors
import safetensors.torch
import torch
import tqdm
from collections import OrderedDict
parser = argparse.ArgumentParser()
parser.add_argument("diffusers_path", type=str,
help="Path to the original Flux diffusers folder.")
parser.add_argument("quantized_state_dict_path", type=str,
help="Path to the ComfyUI all in one template file.")
parser.add_argument("flux_path", type=str,
help="Output path for the Flux safetensors file.")
parser.add_argument("--do_8_bit", action="store_true",
help="Use 8-bit weights instead of bf16.")
args = parser.parse_args()
flux_path = Path(args.flux_path)
diffusers_path = Path(args.diffusers_path, "transformer")
quantized_state_dict_path = Path(args.quantized_state_dict_path)
do_8_bit = args.do_8_bit
if not os.path.exists(flux_path.parent):
os.makedirs(flux_path.parent)
if not diffusers_path.exists():
print(f"Error: Missing transformer folder: {diffusers_path}")
exit()
original_json_path = Path.joinpath(
diffusers_path, "diffusion_pytorch_model.safetensors.index.json")
if not original_json_path.exists():
print(f"Error: Missing transformer index json: {original_json_path}")
exit()
if not os.path.exists(quantized_state_dict_path):
print(
f"Error: Missing quantized state dict file: {args.quantized_state_dict_path}")
exit()
with open(original_json_path, "r", encoding="utf-8") as f:
original_json = json.load(f)
diffusers_map = {
"time_in.in_layer.weight": [
"time_text_embed.timestep_embedder.linear_1.weight",
],
"time_in.in_layer.bias": [
"time_text_embed.timestep_embedder.linear_1.bias",
],
"time_in.out_layer.weight": [
"time_text_embed.timestep_embedder.linear_2.weight",
],
"time_in.out_layer.bias": [
"time_text_embed.timestep_embedder.linear_2.bias",
],
"vector_in.in_layer.weight": [
"time_text_embed.text_embedder.linear_1.weight",
],
"vector_in.in_layer.bias": [
"time_text_embed.text_embedder.linear_1.bias",
],
"vector_in.out_layer.weight": [
"time_text_embed.text_embedder.linear_2.weight",
],
"vector_in.out_layer.bias": [
"time_text_embed.text_embedder.linear_2.bias",
],
"guidance_in.in_layer.weight": [
"time_text_embed.guidance_embedder.linear_1.weight",
],
"guidance_in.in_layer.bias": [
"time_text_embed.guidance_embedder.linear_1.bias",
],
"guidance_in.out_layer.weight": [
"time_text_embed.guidance_embedder.linear_2.weight",
],
"guidance_in.out_layer.bias": [
"time_text_embed.guidance_embedder.linear_2.bias",
],
"txt_in.weight": [
"context_embedder.weight",
],
"txt_in.bias": [
"context_embedder.bias",
],
"img_in.weight": [
"x_embedder.weight",
],
"img_in.bias": [
"x_embedder.bias",
],
"double_blocks.().img_mod.lin.weight": [
"norm1.linear.weight",
],
"double_blocks.().img_mod.lin.bias": [
"norm1.linear.bias",
],
"double_blocks.().txt_mod.lin.weight": [
"norm1_context.linear.weight",
],
"double_blocks.().txt_mod.lin.bias": [
"norm1_context.linear.bias",
],
"double_blocks.().img_attn.qkv.weight": [
"attn.to_q.weight",
"attn.to_k.weight",
"attn.to_v.weight",
],
"double_blocks.().img_attn.qkv.bias": [
"attn.to_q.bias",
"attn.to_k.bias",
"attn.to_v.bias",
],
"double_blocks.().txt_attn.qkv.weight": [
"attn.add_q_proj.weight",
"attn.add_k_proj.weight",
"attn.add_v_proj.weight",
],
"double_blocks.().txt_attn.qkv.bias": [
"attn.add_q_proj.bias",
"attn.add_k_proj.bias",
"attn.add_v_proj.bias",
],
"double_blocks.().img_attn.norm.query_norm.scale": [
"attn.norm_q.weight",
],
"double_blocks.().img_attn.norm.key_norm.scale": [
"attn.norm_k.weight",
],
"double_blocks.().txt_attn.norm.query_norm.scale": [
"attn.norm_added_q.weight",
],
"double_blocks.().txt_attn.norm.key_norm.scale": [
"attn.norm_added_k.weight",
],
"double_blocks.().img_mlp.0.weight": [
"ff.net.0.proj.weight",
],
"double_blocks.().img_mlp.0.bias": [
"ff.net.0.proj.bias",
],
"double_blocks.().img_mlp.2.weight": [
"ff.net.2.weight",
],
"double_blocks.().img_mlp.2.bias": [
"ff.net.2.bias",
],
"double_blocks.().txt_mlp.0.weight": [
"ff_context.net.0.proj.weight",
],
"double_blocks.().txt_mlp.0.bias": [
"ff_context.net.0.proj.bias",
],
"double_blocks.().txt_mlp.2.weight": [
"ff_context.net.2.weight",
],
"double_blocks.().txt_mlp.2.bias": [
"ff_context.net.2.bias",
],
"double_blocks.().img_attn.proj.weight": [
"attn.to_out.0.weight",
],
"double_blocks.().img_attn.proj.bias": [
"attn.to_out.0.bias",
],
"double_blocks.().txt_attn.proj.weight": [
"attn.to_add_out.weight",
],
"double_blocks.().txt_attn.proj.bias": [
"attn.to_add_out.bias",
],
"single_blocks.().modulation.lin.weight": [
"norm.linear.weight",
],
"single_blocks.().modulation.lin.bias": [
"norm.linear.bias",
],
"single_blocks.().linear1.weight": [
"attn.to_q.weight",
"attn.to_k.weight",
"attn.to_v.weight",
"proj_mlp.weight",
],
"single_blocks.().linear1.bias": [
"attn.to_q.bias",
"attn.to_k.bias",
"attn.to_v.bias",
"proj_mlp.bias",
],
"single_blocks.().linear2.weight": [
"proj_out.weight",
],
"single_blocks.().norm.query_norm.scale": [
"attn.norm_q.weight",
],
"single_blocks.().norm.key_norm.scale": [
"attn.norm_k.weight",
],
"single_blocks.().linear2.weight": [
"proj_out.weight",
],
"single_blocks.().linear2.bias": [
"proj_out.bias",
],
"final_layer.linear.weight": [
"proj_out.weight",
],
"final_layer.linear.bias": [
"proj_out.bias",
],
"final_layer.adaLN_modulation.1.weight": [
"norm_out.linear.weight",
],
"final_layer.adaLN_modulation.1.bias": [
"norm_out.linear.bias",
],
}
def is_in_diffusers_map(k):
for values in diffusers_map.values():
for value in values:
if k.endswith(value):
return True
return False
diffusers = {k: Path.joinpath(diffusers_path, v)
for k, v in original_json["weight_map"].items() if is_in_diffusers_map(k)}
original_safetensors = set(diffusers.values())
# determine the number of transformer blocks
transformer_blocks = 0
single_transformer_blocks = 0
for key in diffusers.keys():
print(key)
if key.startswith("transformer_blocks."):
print(key)
block = int(key.split(".")[1])
if block >= transformer_blocks:
transformer_blocks = block + 1
elif key.startswith("single_transformer_blocks."):
block = int(key.split(".")[1])
if block >= single_transformer_blocks:
single_transformer_blocks = block + 1
print(f"Transformer blocks: {transformer_blocks}")
print(f"Single transformer blocks: {single_transformer_blocks}")
for file in original_safetensors:
if not file.exists():
print(f"Error: Missing transformer safetensors file: {file}")
exit()
original_safetensors = {f: safetensors.safe_open(
f, framework="pt", device="cpu") for f in original_safetensors}
def swap_scale_shift(weight):
shift, scale = weight.chunk(2, dim=0)
new_weight = torch.cat([scale, shift], dim=0)
return new_weight
flux_values = {}
for b in range(transformer_blocks):
for key, weights in diffusers_map.items():
if key.startswith("double_blocks."):
block_prefix = f"transformer_blocks.{b}."
found = True
for weight in weights:
if not (f"{block_prefix}{weight}" in diffusers):
found = False
if found:
flux_values[key.replace("()", f"{b}")] = [
f"{block_prefix}{weight}" for weight in weights]
for b in range(single_transformer_blocks):
for key, weights in diffusers_map.items():
if key.startswith("single_blocks."):
block_prefix = f"single_transformer_blocks.{b}."
found = True
for weight in weights:
if not (f"{block_prefix}{weight}" in diffusers):
found = False
if found:
flux_values[key.replace("()", f"{b}")] = [
f"{block_prefix}{weight}" for weight in weights]
for key, weights in diffusers_map.items():
if not (key.startswith("double_blocks.") or key.startswith("single_blocks.")):
found = True
for weight in weights:
if not (f"{weight}" in diffusers):
found = False
if found:
flux_values[key] = [f"{weight}" for weight in weights]
flux = {}
for key, values in tqdm.tqdm(flux_values.items()):
if len(values) == 1:
flux[key] = original_safetensors[diffusers[values[0]]
].get_tensor(values[0]).to("cpu")
else:
flux[key] = torch.cat(
[
original_safetensors[diffusers[value]
].get_tensor(value).to("cpu")
for value in values
]
)
if "norm_out.linear.weight" in diffusers:
flux["final_layer.adaLN_modulation.1.weight"] = swap_scale_shift(
original_safetensors[diffusers["norm_out.linear.weight"]].get_tensor(
"norm_out.linear.weight").to("cpu")
)
if "norm_out.linear.bias" in diffusers:
flux["final_layer.adaLN_modulation.1.bias"] = swap_scale_shift(
original_safetensors[diffusers["norm_out.linear.bias"]].get_tensor(
"norm_out.linear.bias").to("cpu")
)
def stochastic_round_to(tensor, dtype=torch.float8_e4m3fn):
# Define the float8 range
min_val = torch.finfo(dtype).min
max_val = torch.finfo(dtype).max
# Clip values to float8 range
tensor = torch.clamp(tensor, min_val, max_val)
# Convert to float32 for calculations
tensor = tensor.float()
# Get the nearest representable float8 values
lower = torch.floor(tensor * 256) / 256
upper = torch.ceil(tensor * 256) / 256
# Calculate the probability of rounding up
prob = (tensor - lower) / (upper - lower)
# Generate random values for stochastic rounding
rand = torch.rand_like(tensor)
# Perform stochastic rounding
rounded = torch.where(rand < prob, upper, lower)
# Convert back to float8
return rounded.to(dtype)
# set all the keys to bf16
for key in flux.keys():
if do_8_bit:
flux[key] = stochastic_round_to(
flux[key], torch.float8_e4m3fn).to('cpu')
else:
flux[key] = flux[key].clone().to('cpu', torch.bfloat16)
# load the quantized state dict
quantized_state_dict = safetensors.torch.load_file(quantized_state_dict_path)
transformer_pre = "model.diffusion_model."
did_print = False
# remove old parts
for key in list(quantized_state_dict.keys()):
if key.startswith(transformer_pre):
if not did_print:
# print("dtype: ", quantized_state_dict[key].dtype)
did_print = True
del quantized_state_dict[key]
# add the new parts
for key, value in flux.items():
quantized_state_dict[transformer_pre + key] = value
meta = OrderedDict()
meta['format'] = 'pt'
# date format like 2024-08-01 YYYY-MM-DD
meta['modelspec.date'] = date.today().strftime("%Y-%m-%d")
meta['modelspec.title'] = "Flex.1-alpha"
meta['modelspec.author'] = "Ostris, LLC"
meta['modelspec.license'] = "Apache-2.0"
meta['modelspec.implementation'] = "https://github.com/black-forest-labs/flux"
meta['modelspec.architecture'] = "Flex.1-alpha"
meta['modelspec.description'] = "Flex.1-alpha"
os.makedirs(os.path.dirname(flux_path), exist_ok=True)
print(f"Saving to {flux_path}")
safetensors.torch.save_file(quantized_state_dict, flux_path, metadata=meta)
print("Done.")

View File

@@ -0,0 +1,91 @@
# currently only works with flux as support is not quite there yet
import argparse
import os.path
from collections import OrderedDict
parser = argparse.ArgumentParser()
parser.add_argument(
'input_path',
type=str,
help='Path to original sdxl model'
)
parser.add_argument(
'output_path',
type=str,
help='output path'
)
args = parser.parse_args()
args.input_path = os.path.abspath(args.input_path)
args.output_path = os.path.abspath(args.output_path)
from safetensors.torch import load_file, save_file
meta = OrderedDict()
meta['format'] = 'pt'
state_dict = load_file(args.input_path)
# peft doesnt have an alpha so we need to scale the weights
alpha_keys = [
'lora_transformer_single_transformer_blocks_0_attn_to_q.alpha' # flux
]
# keys where the rank is in the first dimension
rank_idx0_keys = [
'lora_transformer_single_transformer_blocks_0_attn_to_q.lora_down.weight'
# 'transformer.single_transformer_blocks.0.attn.to_q.lora_A.weight'
]
alpha = None
rank = None
for key in rank_idx0_keys:
if key in state_dict:
rank = int(state_dict[key].shape[0])
break
if rank is None:
raise ValueError(f'Could not find rank in state dict')
for key in alpha_keys:
if key in state_dict:
alpha = int(state_dict[key])
break
if alpha is None:
# set to rank if not found
alpha = rank
up_multiplier = alpha / rank
new_state_dict = {}
for key, value in state_dict.items():
if key.endswith('.alpha'):
continue
orig_dtype = value.dtype
new_val = value.float() * up_multiplier
new_key = key
new_key = new_key.replace('lora_transformer_', 'transformer.')
for i in range(100):
new_key = new_key.replace(f'transformer_blocks_{i}_', f'transformer_blocks.{i}.')
new_key = new_key.replace('lora_down', 'lora_A')
new_key = new_key.replace('lora_up', 'lora_B')
new_key = new_key.replace('_lora', '.lora')
new_key = new_key.replace('attn_', 'attn.')
new_key = new_key.replace('ff_', 'ff.')
new_key = new_key.replace('context_net_', 'context.net.')
new_key = new_key.replace('0_proj', '0.proj')
new_key = new_key.replace('norm_linear', 'norm.linear')
new_key = new_key.replace('norm_out_linear', 'norm_out.linear')
new_key = new_key.replace('to_out_', 'to_out.')
new_state_dict[new_key] = new_val.to(orig_dtype)
save_file(new_state_dict, args.output_path, meta)
print(f'Saved to {args.output_path}')

View File

@@ -0,0 +1,245 @@
import os
from tqdm import tqdm
import argparse
from collections import OrderedDict
parser = argparse.ArgumentParser(description="Extract LoRA from Flex")
parser.add_argument("--base", type=str, default="ostris/Flex.1-alpha", help="Base model path")
parser.add_argument("--tuned", type=str, required=True, help="Tuned model path")
parser.add_argument("--output", type=str, required=True, help="Output path for lora")
parser.add_argument("--rank", type=int, default=32, help="LoRA rank for extraction")
parser.add_argument("--gpu", type=int, default=0, help="GPU to process extraction")
parser.add_argument("--full", action="store_true", help="Do a full transformer extraction, not just transformer blocks")
args = parser.parse_args()
if True:
# set cuda environment variable
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
import torch
from safetensors.torch import load_file, save_file
from lycoris.utils import extract_linear, extract_conv, make_sparse
from diffusers import FluxTransformer2DModel
base = args.base
tuned = args.tuned
output_path = args.output
dim = args.rank
os.makedirs(os.path.dirname(output_path), exist_ok=True)
state_dict_base = {}
state_dict_tuned = {}
output_dict = {}
@torch.no_grad()
def extract_diff(
base_unet,
db_unet,
mode="fixed",
linear_mode_param=0,
conv_mode_param=0,
extract_device="cpu",
use_bias=False,
sparsity=0.98,
# small_conv=True,
small_conv=False,
):
UNET_TARGET_REPLACE_MODULE = [
"Linear",
"Conv2d",
"LayerNorm",
"GroupNorm",
"GroupNorm32",
"LoRACompatibleLinear",
"LoRACompatibleConv"
]
LORA_PREFIX_UNET = "transformer"
def make_state_dict(
prefix,
root_module: torch.nn.Module,
target_module: torch.nn.Module,
target_replace_modules,
):
loras = {}
temp = {}
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
temp[name] = module
for name, module in tqdm(
list((n, m) for n, m in target_module.named_modules() if n in temp)
):
weights = temp[name]
lora_name = prefix + "." + name
# lora_name = lora_name.replace(".", "_")
layer = module.__class__.__name__
if 'transformer_blocks' not in lora_name and not args.full:
continue
if layer in {
"Linear",
"Conv2d",
"LayerNorm",
"GroupNorm",
"GroupNorm32",
"Embedding",
"LoRACompatibleLinear",
"LoRACompatibleConv"
}:
root_weight = module.weight
try:
if torch.allclose(root_weight, weights.weight):
continue
except:
continue
else:
continue
module = module.to(extract_device, torch.float32)
weights = weights.to(extract_device, torch.float32)
if mode == "full":
decompose_mode = "full"
elif layer == "Linear":
weight, decompose_mode = extract_linear(
(root_weight - weights.weight),
mode,
linear_mode_param,
device=extract_device,
)
if decompose_mode == "low rank":
extract_a, extract_b, diff = weight
elif layer == "Conv2d":
is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1
weight, decompose_mode = extract_conv(
(root_weight - weights.weight),
mode,
linear_mode_param if is_linear else conv_mode_param,
device=extract_device,
)
if decompose_mode == "low rank":
extract_a, extract_b, diff = weight
if small_conv and not is_linear and decompose_mode == "low rank":
dim = extract_a.size(0)
(extract_c, extract_a, _), _ = extract_conv(
extract_a.transpose(0, 1),
"fixed",
dim,
extract_device,
True,
)
extract_a = extract_a.transpose(0, 1)
extract_c = extract_c.transpose(0, 1)
loras[f"{lora_name}.lora_mid.weight"] = (
extract_c.detach().cpu().contiguous().half()
)
diff = (
(
root_weight
- torch.einsum(
"i j k l, j r, p i -> p r k l",
extract_c,
extract_a.flatten(1, -1),
extract_b.flatten(1, -1),
)
)
.detach()
.cpu()
.contiguous()
)
del extract_c
else:
module = module.to("cpu")
weights = weights.to("cpu")
continue
if decompose_mode == "low rank":
loras[f"{lora_name}.lora_A.weight"] = (
extract_a.detach().cpu().contiguous().half()
)
loras[f"{lora_name}.lora_B.weight"] = (
extract_b.detach().cpu().contiguous().half()
)
# loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half()
if use_bias:
diff = diff.detach().cpu().reshape(extract_b.size(0), -1)
sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce()
indices = sparse_diff.indices().to(torch.int16)
values = sparse_diff.values().half()
loras[f"{lora_name}.bias_indices"] = indices
loras[f"{lora_name}.bias_values"] = values
loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to(
torch.int16
)
del extract_a, extract_b, diff
elif decompose_mode == "full":
if "Norm" in layer:
w_key = "w_norm"
b_key = "b_norm"
else:
w_key = "diff"
b_key = "diff_b"
weight_diff = module.weight - weights.weight
loras[f"{lora_name}.{w_key}"] = (
weight_diff.detach().cpu().contiguous().half()
)
if getattr(weights, "bias", None) is not None:
bias_diff = module.bias - weights.bias
loras[f"{lora_name}.{b_key}"] = (
bias_diff.detach().cpu().contiguous().half()
)
else:
raise NotImplementedError
module = module.to("cpu", torch.bfloat16)
weights = weights.to("cpu", torch.bfloat16)
return loras
all_loras = {}
all_loras |= make_state_dict(
LORA_PREFIX_UNET,
base_unet,
db_unet,
UNET_TARGET_REPLACE_MODULE,
)
del base_unet, db_unet
if torch.cuda.is_available():
torch.cuda.empty_cache()
all_lora_name = set()
for k in all_loras:
lora_name, weight = k.rsplit(".", 1)
all_lora_name.add(lora_name)
print(len(all_lora_name))
return all_loras
# find all the .safetensors files and load them
print("Loading Base")
base_model = FluxTransformer2DModel.from_pretrained(base, subfolder="transformer", torch_dtype=torch.bfloat16)
print("Loading Tuned")
tuned_model = FluxTransformer2DModel.from_pretrained(tuned, subfolder="transformer", torch_dtype=torch.bfloat16)
output_dict = extract_diff(
base_model,
tuned_model,
mode="fixed",
linear_mode_param=dim,
conv_mode_param=dim,
extract_device="cuda",
use_bias=False,
sparsity=0.98,
small_conv=False,
)
meta = OrderedDict()
meta['format'] = 'pt'
save_file(output_dict, output_path, metadata=meta)
print("Done")

View File

@@ -0,0 +1,20 @@
import argparse
import torch
import os
from diffusers import StableDiffusionPipeline
import sys
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
# add project root to path
sys.path.append(PROJECT_ROOT)
SAMPLER_SCALES_ROOT = os.path.join(PROJECT_ROOT, 'toolkit', 'samplers_scales')
parser = argparse.ArgumentParser(description='Process some images.')
add_arg = parser.add_argument
add_arg('--model', type=str, required=True, help='Path to model')
add_arg('--sampler', type=str, required=True, help='Name of sampler')
args = parser.parse_args()

View File

@@ -1,5 +1,9 @@
import argparse
from collections import OrderedDict
import sys
import os
ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.append(ROOT_DIR)
import torch

View File

@@ -0,0 +1,42 @@
import torch
from safetensors.torch import save_file, load_file
from collections import OrderedDict
meta = OrderedDict()
meta["format"] ="pt"
attn_dict = load_file("/mnt/Train/out/ip_adapter/sd15_bigG/sd15_bigG_000266000.safetensors")
state_dict = load_file("/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors")
attn_list = []
for key, value in state_dict.items():
if "attn1" in key:
attn_list.append(key)
attn_names = ['down_blocks.0.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.0.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'down_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.1.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.2.attentions.2.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.0.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.1.transformer_blocks.0.attn2.processor', 'up_blocks.3.attentions.2.transformer_blocks.0.attn2.processor', 'mid_block.attentions.0.transformer_blocks.0.attn2.processor']
adapter_names = []
for i in range(100):
if f'te_adapter.adapter_modules.{i}.to_k_adapter.weight' in attn_dict:
adapter_names.append(f"te_adapter.adapter_modules.{i}.adapter")
for i in range(len(adapter_names)):
adapter_name = adapter_names[i]
attn_name = attn_names[i]
adapter_k_name = adapter_name[:-8] + '.to_k_adapter.weight'
adapter_v_name = adapter_name[:-8] + '.to_v_adapter.weight'
state_k_name = attn_name.replace(".processor", ".to_k.weight")
state_v_name = attn_name.replace(".processor", ".to_v.weight")
if adapter_k_name in attn_dict:
state_dict[state_k_name] = attn_dict[adapter_k_name]
state_dict[state_v_name] = attn_dict[adapter_v_name]
else:
print("adapter_k_name", adapter_k_name)
print("state_k_name", state_k_name)
for key, value in state_dict.items():
state_dict[key] = value.cpu().to(torch.float16)
save_file(state_dict, "/home/jaret/Dev/models/hf/OstrisDiffusionV1/unet/diffusion_pytorch_model.safetensors", metadata=meta)
print("Done")

View File

@@ -0,0 +1,65 @@
import argparse
from PIL import Image
from PIL.ImageOps import exif_transpose
from tqdm import tqdm
import os
parser = argparse.ArgumentParser(description='Process some images.')
parser.add_argument("input_folder", type=str, help="Path to folder containing images")
args = parser.parse_args()
img_types = ['.jpg', '.jpeg', '.png', '.webp']
# find all images in the input folder
images = []
for root, _, files in os.walk(args.input_folder):
for file in files:
if file.lower().endswith(tuple(img_types)):
images.append(os.path.join(root, file))
print(f"Found {len(images)} images")
num_skipped = 0
num_repaired = 0
num_deleted = 0
pbar = tqdm(total=len(images), desc=f"Repaired {num_repaired} images", unit="image")
for img_path in images:
filename = os.path.basename(img_path)
filename_no_ext, file_extension = os.path.splitext(filename)
# if it is jpg, ignore
if file_extension.lower() == '.jpg':
num_skipped += 1
pbar.update(1)
continue
try:
img = Image.open(img_path)
except Exception as e:
print(f"Error opening {img_path}: {e}")
# delete it
os.remove(img_path)
num_deleted += 1
pbar.update(1)
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
continue
try:
img = exif_transpose(img)
except Exception as e:
print(f"Error rotating {img_path}: {e}")
new_path = os.path.join(os.path.dirname(img_path), filename_no_ext + '.jpg')
img = img.convert("RGB")
img.save(new_path, quality=95)
# remove the old file
os.remove(img_path)
num_repaired += 1
pbar.update(1)
# update pbar
pbar.set_description(f"Repaired {num_repaired} images, Skipped {num_skipped}, Deleted {num_deleted}")
print("Done")

309
scripts/update_sponsors.py Normal file
View File

@@ -0,0 +1,309 @@
import os
import requests
import json
from datetime import datetime
from dotenv import load_dotenv
# Load environment variables from .env file
env_path = os.path.join(os.path.dirname(os.path.dirname(__file__)), ".env")
load_dotenv(dotenv_path=env_path)
# API credentials
PATREON_TOKEN = os.getenv("PATREON_ACCESS_TOKEN")
GITHUB_TOKEN = os.getenv("GITHUB_TOKEN")
GITHUB_USERNAME = os.getenv("GITHUB_USERNAME")
GITHUB_ORG = os.getenv("GITHUB_ORG") # Organization name (optional)
# Output file
README_PATH = "SUPPORTERS.md"
def fetch_patreon_supporters():
"""Fetch current Patreon supporters"""
print("Fetching Patreon supporters...")
headers = {
"Authorization": f"Bearer {PATREON_TOKEN}",
"Content-Type": "application/json"
}
url = "https://www.patreon.com/api/oauth2/v2/campaigns"
try:
# First get the campaign ID
campaign_response = requests.get(url, headers=headers)
campaign_response.raise_for_status()
campaign_data = campaign_response.json()
if not campaign_data.get('data'):
print("No campaigns found for this Patreon account")
return []
campaign_id = campaign_data['data'][0]['id']
# Now get the supporters for this campaign
members_url = f"https://www.patreon.com/api/oauth2/v2/campaigns/{campaign_id}/members"
params = {
"include": "user",
"fields[member]": "full_name,is_follower,patron_status", # Removed profile_url
"fields[user]": "image_url"
}
supporters = []
while members_url:
members_response = requests.get(members_url, headers=headers, params=params)
members_response.raise_for_status()
members_data = members_response.json()
# Process the response to extract active patrons
for member in members_data.get('data', []):
attributes = member.get('attributes', {})
# Only include active patrons
if attributes.get('patron_status') == 'active_patron':
name = attributes.get('full_name', 'Anonymous Supporter')
# Get user data which contains the profile image
user_id = member.get('relationships', {}).get('user', {}).get('data', {}).get('id')
profile_image = None
profile_url = None # Removed profile_url since it's not supported
if user_id:
for included in members_data.get('included', []):
if included.get('id') == user_id and included.get('type') == 'user':
profile_image = included.get('attributes', {}).get('image_url')
break
supporters.append({
'name': name,
'profile_image': profile_image,
'profile_url': profile_url, # This will be None
'platform': 'Patreon',
'amount': 0 # Placeholder, as Patreon API doesn't provide this in the current response
})
# Handle pagination
members_url = members_data.get('links', {}).get('next')
print(f"Found {len(supporters)} active Patreon supporters")
return supporters
except requests.exceptions.RequestException as e:
print(f"Error fetching Patreon data: {e}")
print(f"Response content: {e.response.content if hasattr(e, 'response') else 'No response content'}")
return []
def fetch_github_sponsors():
"""Fetch current GitHub sponsors for a user or organization"""
print("Fetching GitHub sponsors...")
headers = {
"Authorization": f"Bearer {GITHUB_TOKEN}",
"Accept": "application/vnd.github.v3+json"
}
# Determine if we're fetching for a user or an organization
entity_type = "organization" if GITHUB_ORG else "user"
entity_name = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
if not entity_name:
print("Error: Neither GITHUB_USERNAME nor GITHUB_ORG is set")
return []
# Different GraphQL query structure based on entity type
if entity_type == "user":
query = """
query {
user(login: "%s") {
sponsorshipsAsMaintainer(first: 100) {
nodes {
sponsorEntity {
... on User {
login
name
avatarUrl
url
}
... on Organization {
login
name
avatarUrl
url
}
}
tier {
monthlyPriceInDollars
}
isOneTimePayment
isActive
}
}
}
}
""" % entity_name
else: # organization
query = """
query {
organization(login: "%s") {
sponsorshipsAsMaintainer(first: 100) {
nodes {
sponsorEntity {
... on User {
login
name
avatarUrl
url
}
... on Organization {
login
name
avatarUrl
url
}
}
tier {
monthlyPriceInDollars
}
isOneTimePayment
isActive
}
}
}
}
""" % entity_name
try:
response = requests.post(
"https://api.github.com/graphql",
headers=headers,
json={"query": query}
)
response.raise_for_status()
data = response.json()
# Process the response - the path to the data differs based on entity type
if entity_type == "user":
sponsors_data = data.get('data', {}).get('user', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
else:
sponsors_data = data.get('data', {}).get('organization', {}).get('sponsorshipsAsMaintainer', {}).get('nodes', [])
sponsors = []
for sponsor in sponsors_data:
# Only include active sponsors
if sponsor.get('isActive'):
entity = sponsor.get('sponsorEntity', {})
name = entity.get('name') or entity.get('login', 'Anonymous Sponsor')
profile_image = entity.get('avatarUrl')
profile_url = entity.get('url')
amount = sponsor.get('tier', {}).get('monthlyPriceInDollars', 0)
sponsors.append({
'name': name,
'profile_image': profile_image,
'profile_url': profile_url,
'platform': 'GitHub Sponsors',
'amount': amount
})
print(f"Found {len(sponsors)} active GitHub sponsors for {entity_type} '{entity_name}'")
return sponsors
except requests.exceptions.RequestException as e:
print(f"Error fetching GitHub sponsors data: {e}")
return []
def generate_readme(supporters):
"""Generate a README.md file with supporter information"""
print(f"Generating {README_PATH}...")
# Sort supporters by amount (descending) and then by name
supporters.sort(key=lambda x: (-x['amount'], x['name'].lower()))
# Determine the proper footer links based on what's configured
github_entity = GITHUB_ORG if GITHUB_ORG else GITHUB_USERNAME
github_entity_type = "orgs" if GITHUB_ORG else "sponsors"
github_sponsor_url = f"https://github.com/{github_entity_type}/{github_entity}"
with open(README_PATH, "w", encoding="utf-8") as f:
f.write("## Support My Work\n\n")
f.write("If you enjoy my work, or use it for commercial purposes, please consider sponsoring me so I can continue to maintain it. Every bit helps! \n\n")
# Create appropriate call-to-action based on what's configured
cta_parts = []
if github_entity:
cta_parts.append(f"[Become a sponsor on GitHub]({github_sponsor_url})")
if PATREON_TOKEN:
cta_parts.append("[support me on Patreon](https://www.patreon.com/ostris)")
if cta_parts:
if GITHUB_ORG:
f.write(f"{' or '.join(cta_parts)}.\n\n")
f.write("Thank you to all my current supporters!\n\n")
f.write(f"_Last updated: {datetime.now().strftime('%Y-%m-%d')}_\n\n")
# Write GitHub Sponsors section
github_sponsors = [s for s in supporters if s['platform'] == 'GitHub Sponsors']
if github_sponsors:
f.write("### GitHub Sponsors\n\n")
for sponsor in github_sponsors:
if sponsor['profile_image']:
f.write(f"<a href=\"{sponsor['profile_url']}\" title=\"{sponsor['name']}\"><img src=\"{sponsor['profile_image']}\" width=\"50\" height=\"50\" alt=\"{sponsor['name']}\" style=\"border-radius:50%\"></a> ")
else:
f.write(f"[{sponsor['name']}]({sponsor['profile_url']}) ")
f.write("\n\n")
# Write Patreon section
patreon_supporters = [s for s in supporters if s['platform'] == 'Patreon']
if patreon_supporters:
f.write("### Patreon Supporters\n\n")
for supporter in patreon_supporters:
if supporter['profile_image']:
f.write(f"<a href=\"{supporter['profile_url']}\" title=\"{supporter['name']}\"><img src=\"{supporter['profile_image']}\" width=\"50\" height=\"50\" alt=\"{supporter['name']}\" style=\"border-radius:50%\"></a> ")
else:
f.write(f"[{supporter['name']}]({supporter['profile_url']}) ")
f.write("\n\n")
f.write("\n---\n\n")
print(f"Successfully generated {README_PATH} with {len(supporters)} supporters!")
def main():
"""Main function"""
print("Starting supporter data collection...")
# Check if required environment variables are set
missing_vars = []
if not GITHUB_TOKEN:
missing_vars.append("GITHUB_TOKEN")
# Either username or org is required for GitHub
if not GITHUB_USERNAME and not GITHUB_ORG:
missing_vars.append("GITHUB_USERNAME or GITHUB_ORG")
# Patreon token is optional but warn if missing
patreon_enabled = bool(PATREON_TOKEN)
if missing_vars:
print(f"Error: Missing required environment variables: {', '.join(missing_vars)}")
print("Please add them to your .env file")
return
if not patreon_enabled:
print("Warning: PATREON_ACCESS_TOKEN not set. Will only fetch GitHub sponsors.")
# Fetch data from both platforms
patreon_supporters = fetch_patreon_supporters() if PATREON_TOKEN else []
github_sponsors = fetch_github_sponsors()
# Combine supporters from both platforms
all_supporters = patreon_supporters + github_sponsors
if not all_supporters:
print("No supporters found on either platform")
return
# Generate README
generate_readme(all_supporters)
if __name__ == "__main__":
main()

View File

@@ -54,6 +54,7 @@ parser.add_argument('--name', type=str, default='stable_diffusion', help='name f
parser.add_argument('--sdxl', action='store_true', help='is sdxl model')
parser.add_argument('--refiner', action='store_true', help='is refiner model')
parser.add_argument('--ssd', action='store_true', help='is ssd model')
parser.add_argument('--vega', action='store_true', help='is vega model')
parser.add_argument('--sd2', action='store_true', help='is sd 2 model')
args = parser.parse_args()
@@ -66,15 +67,15 @@ print(f'Loading diffusers model')
ignore_ldm_begins_with = []
diffusers_file_path = file_path
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
if args.ssd:
diffusers_file_path = "segmind/SSD-1B"
if args.vega:
diffusers_file_path = "segmind/Segmind-Vega"
# if args.refiner:
# diffusers_file_path = "stabilityai/stable-diffusion-xl-refiner-1.0"
diffusers_file_path = file_path if len(args.file_1) == 1 else args.file_1[1]
if not args.refiner:
diffusers_model_config = ModelConfig(
@@ -82,6 +83,7 @@ if not args.refiner:
is_xl=args.sdxl,
is_v2=args.sd2,
is_ssd=args.ssd,
is_vega=args.vega,
dtype=dtype,
)
diffusers_sd = StableDiffusion(
@@ -157,7 +159,7 @@ te_suffix = ''
proj_pattern_weight = None
proj_pattern_bias = None
text_proj_layer = None
if args.sdxl or args.ssd:
if args.sdxl or args.ssd or args.vega:
te_suffix = '1'
ldm_res_block_prefix = "conditioner.embedders.1.model.transformer.resblocks"
proj_pattern_weight = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
@@ -176,10 +178,13 @@ if args.sd2:
proj_pattern_bias = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
text_proj_layer = "cond_stage_model.model.text_projection"
if args.sdxl or args.sd2 or args.ssd or args.refiner:
if args.sdxl or args.sd2 or args.ssd or args.refiner or args.vega:
if "conditioner.embedders.1.model.text_projection" in ldm_dict_keys:
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection"].shape[0])
elif "conditioner.embedders.1.model.text_projection.weight" in ldm_dict_keys:
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
d_model = int(ldm_state_dict["conditioner.embedders.1.model.text_projection.weight"].shape[0])
elif "conditioner.embedders.0.model.text_projection" in ldm_dict_keys:
# d_model = int(checkpoint[prefix + "text_projection"].shape[0]))
d_model = int(ldm_state_dict["conditioner.embedders.0.model.text_projection"].shape[0])
@@ -191,6 +196,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
try:
match = re.match(proj_pattern_weight, ldm_key)
if match:
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight":
print("here")
number = int(match.group(1))
new_val = torch.cat([
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"],
@@ -217,6 +224,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
],
}
matched_ldm_keys.append(ldm_key)
# text_model_dict[new_key + ".q_proj.weight"] = checkpoint[key][:d_model, :]
# text_model_dict[new_key + ".k_proj.weight"] = checkpoint[key][d_model: d_model * 2, :]
# text_model_dict[new_key + ".v_proj.weight"] = checkpoint[key][d_model * 2:, :]
@@ -266,6 +275,8 @@ if args.sdxl or args.sd2 or args.ssd or args.refiner:
],
}
matched_ldm_keys.append(ldm_key)
# add diffusers operators
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias"] = {
"slice": [
@@ -298,6 +309,9 @@ for ldm_key in ldm_dict_keys:
ldm_shape_tuple = ldm_state_dict[ldm_key].shape
ldm_reduced_shape_tuple = get_reduced_shape(ldm_shape_tuple)
for diffusers_key in diffusers_dict_keys:
if ldm_key == "conditioner.embedders.1.model.transformer.resblocks.0.attn.in_proj_weight" and diffusers_key == "te1_text_model.encoder.layers.0.self_attn.q_proj.weight":
print("here")
diffusers_shape_tuple = diffusers_state_dict[diffusers_key].shape
diffusers_reduced_shape_tuple = get_reduced_shape(diffusers_shape_tuple)
@@ -356,6 +370,8 @@ if args.sdxl:
name += '_sdxl'
elif args.ssd:
name += '_ssd'
elif args.vega:
name += '_vega'
elif args.refiner:
name += '_refiner'
elif args.sd2:

View File

@@ -0,0 +1,180 @@
import os
import torch
from transformers import T5EncoderModel, T5Tokenizer
from diffusers import StableDiffusionPipeline, UNet2DConditionModel, PixArtSigmaPipeline, Transformer2DModel, PixArtTransformer2DModel
from safetensors.torch import load_file, save_file
from collections import OrderedDict
import json
# model_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_v01_000527000"
# te_path = "google/flan-t5-xl"
# te_aug_path = "/mnt/Train/out/ip_adapter/t5xx_sd15_v1/t5xx_sd15_v1_000032000.safetensors"
# output_path = "/home/jaret/Dev/models/hf/kl-f16-d42_sd15_t5xl_raw"
model_path = "/home/jaret/Dev/models/hf/objective-reality-16ch"
te_path = "google/flan-t5-xl"
te_aug_path = "/mnt/Train2/out/ip_adapter/t5xl-sd15-16ch_v1/t5xl-sd15-16ch_v1_000115000.safetensors"
output_path = "/home/jaret/Dev/models/hf/t5xl-sd15-16ch_sd15_v1"
print("Loading te adapter")
te_aug_sd = load_file(te_aug_path)
print("Loading model")
is_diffusers = (not os.path.exists(model_path)) or os.path.isdir(model_path)
# if "pixart" in model_path.lower():
is_pixart = "pixart" in model_path.lower()
pipeline_class = StableDiffusionPipeline
# transformer = PixArtTransformer2DModel.from_pretrained('PixArt-alpha/PixArt-Sigma-XL-2-512-MS', subfolder='transformer', torch_dtype=torch.float16)
if is_pixart:
pipeline_class = PixArtSigmaPipeline
if is_diffusers:
sd = pipeline_class.from_pretrained(model_path, torch_dtype=torch.float16)
else:
sd = pipeline_class.from_single_file(model_path, torch_dtype=torch.float16)
print("Loading Text Encoder")
# Load the text encoder
te = T5EncoderModel.from_pretrained(te_path, torch_dtype=torch.float16)
# patch it
sd.text_encoder = te
sd.tokenizer = T5Tokenizer.from_pretrained(te_path)
if is_pixart:
unet = sd.transformer
unet_sd = sd.transformer.state_dict()
else:
unet = sd.unet
unet_sd = sd.unet.state_dict()
if is_pixart:
weight_idx = 0
else:
weight_idx = 1
new_cross_attn_dim = None
# count the num of params in state dict
start_params = sum([v.numel() for v in unet_sd.values()])
print("Building")
attn_processor_keys = []
if is_pixart:
transformer: Transformer2DModel = unet
for i, module in transformer.transformer_blocks.named_children():
attn_processor_keys.append(f"transformer_blocks.{i}.attn1")
# cross attention
attn_processor_keys.append(f"transformer_blocks.{i}.attn2")
else:
attn_processor_keys = list(unet.attn_processors.keys())
for name in attn_processor_keys:
cross_attention_dim = None if name.endswith("attn1.processor") or name.endswith("attn.1") or name.endswith(
"attn1") else \
unet.config['cross_attention_dim']
if name.startswith("mid_block"):
hidden_size = unet.config['block_out_channels'][-1]
elif name.startswith("up_blocks"):
block_id = int(name[len("up_blocks.")])
hidden_size = list(reversed(unet.config['block_out_channels']))[block_id]
elif name.startswith("down_blocks"):
block_id = int(name[len("down_blocks.")])
hidden_size = unet.config['block_out_channels'][block_id]
elif name.startswith("transformer"):
hidden_size = unet.config['cross_attention_dim']
else:
# they didnt have this, but would lead to undefined below
raise ValueError(f"unknown attn processor name: {name}")
if cross_attention_dim is None:
pass
else:
layer_name = name.split(".processor")[0]
to_k_adapter = unet_sd[layer_name + ".to_k.weight"]
to_v_adapter = unet_sd[layer_name + ".to_v.weight"]
te_aug_name = None
while True:
if is_pixart:
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
else:
te_aug_name = f"te_adapter.adapter_modules.{weight_idx}.to_k_adapter"
if f"{te_aug_name}.weight" in te_aug_sd:
# increment so we dont redo it next time
weight_idx += 1
break
else:
weight_idx += 1
if weight_idx > 1000:
raise ValueError("Could not find the next weight")
orig_weight_shape_k = list(unet_sd[layer_name + ".to_k.weight"].shape)
new_weight_shape_k = list(te_aug_sd[te_aug_name + ".weight"].shape)
orig_weight_shape_v = list(unet_sd[layer_name + ".to_v.weight"].shape)
new_weight_shape_v = list(te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"].shape)
unet_sd[layer_name + ".to_k.weight"] = te_aug_sd[te_aug_name + ".weight"]
unet_sd[layer_name + ".to_v.weight"] = te_aug_sd[te_aug_name.replace('to_k', 'to_v') + ".weight"]
if new_cross_attn_dim is None:
new_cross_attn_dim = unet_sd[layer_name + ".to_k.weight"].shape[1]
if is_pixart:
# copy the caption_projection weight
del unet_sd['caption_projection.linear_1.bias']
del unet_sd['caption_projection.linear_1.weight']
del unet_sd['caption_projection.linear_2.bias']
del unet_sd['caption_projection.linear_2.weight']
print("Saving unmodified model")
sd = sd.to("cpu", torch.float16)
sd.save_pretrained(
output_path,
safe_serialization=True,
)
# overwrite the unet
if is_pixart:
unet_folder = os.path.join(output_path, "transformer")
else:
unet_folder = os.path.join(output_path, "unet")
# move state_dict to cpu
unet_sd = {k: v.clone().cpu().to(torch.float16) for k, v in unet_sd.items()}
meta = OrderedDict()
meta["format"] = "pt"
print("Patching")
save_file(unet_sd, os.path.join(unet_folder, "diffusion_pytorch_model.safetensors"), meta)
# load the json file
with open(os.path.join(unet_folder, "config.json"), 'r') as f:
config = json.load(f)
config['cross_attention_dim'] = new_cross_attn_dim
if is_pixart:
config['caption_channels'] = None
# save it
with open(os.path.join(unet_folder, "config.json"), 'w') as f:
json.dump(config, f, indent=2)
print("Done")
new_params = sum([v.numel() for v in unet_sd.values()])
# print new and old params with , formatted
print(f"Old params: {start_params:,}")
print(f"New params: {new_params:,}")

62
testing/shrink_pixart.py Normal file
View File

@@ -0,0 +1,62 @@
import torch
from safetensors.torch import load_file, save_file
from collections import OrderedDict
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
state_dict = load_file(model_path)
meta = OrderedDict()
meta["format"] = "pt"
new_state_dict = {}
# Move non-blocks over
for key, value in state_dict.items():
if not key.startswith("transformer_blocks."):
new_state_dict[key] = value
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
'transformer_blocks.{idx}.scale_shift_table']
# New block idx 0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27
current_idx = 0
for i in range(28):
if i not in [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]:
# todo merge in with previous block
for name in block_names:
try:
new_state_dict_key = name.format(idx=current_idx - 1)
old_state_dict_key = name.format(idx=i)
new_state_dict[new_state_dict_key] = (new_state_dict[new_state_dict_key] * 0.5) + (state_dict[old_state_dict_key] * 0.5)
except KeyError:
raise KeyError(f"KeyError: {name.format(idx=current_idx)}")
else:
for name in block_names:
new_state_dict[name.format(idx=current_idx)] = state_dict[name.format(idx=i)]
current_idx += 1
# make sure they are all fp16 and on cpu
for key, value in new_state_dict.items():
new_state_dict[key] = value.to(torch.float16).cpu()
# save the new state dict
save_file(new_state_dict, output_path, metadata=meta)
new_param_count = sum([v.numel() for v in new_state_dict.values()])
old_param_count = sum([v.numel() for v in state_dict.values()])
print(f"Old param count: {old_param_count:,}")
print(f"New param count: {new_param_count:,}")

81
testing/shrink_pixart2.py Normal file
View File

@@ -0,0 +1,81 @@
import torch
from safetensors.torch import load_file, save_file
from collections import OrderedDict
model_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model_orig.safetensors"
output_path = "/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-1024_tiny/transformer/diffusion_pytorch_model.safetensors"
state_dict = load_file(model_path)
meta = OrderedDict()
meta["format"] = "pt"
new_state_dict = {}
# Move non-blocks over
for key, value in state_dict.items():
if not key.startswith("transformer_blocks."):
new_state_dict[key] = value
block_names = ['transformer_blocks.{idx}.attn1.to_k.bias', 'transformer_blocks.{idx}.attn1.to_k.weight',
'transformer_blocks.{idx}.attn1.to_out.0.bias', 'transformer_blocks.{idx}.attn1.to_out.0.weight',
'transformer_blocks.{idx}.attn1.to_q.bias', 'transformer_blocks.{idx}.attn1.to_q.weight',
'transformer_blocks.{idx}.attn1.to_v.bias', 'transformer_blocks.{idx}.attn1.to_v.weight',
'transformer_blocks.{idx}.attn2.to_k.bias', 'transformer_blocks.{idx}.attn2.to_k.weight',
'transformer_blocks.{idx}.attn2.to_out.0.bias', 'transformer_blocks.{idx}.attn2.to_out.0.weight',
'transformer_blocks.{idx}.attn2.to_q.bias', 'transformer_blocks.{idx}.attn2.to_q.weight',
'transformer_blocks.{idx}.attn2.to_v.bias', 'transformer_blocks.{idx}.attn2.to_v.weight',
'transformer_blocks.{idx}.ff.net.0.proj.bias', 'transformer_blocks.{idx}.ff.net.0.proj.weight',
'transformer_blocks.{idx}.ff.net.2.bias', 'transformer_blocks.{idx}.ff.net.2.weight',
'transformer_blocks.{idx}.scale_shift_table']
# Blocks to keep
# keep_blocks = [0, 1, 2, 6, 10, 14, 18, 22, 26, 27]
keep_blocks = [0, 1, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 27]
def weighted_merge(kept_block, removed_block, weight):
return kept_block * (1 - weight) + removed_block * weight
# First, copy all kept blocks to new_state_dict
for i, old_idx in enumerate(keep_blocks):
for name in block_names:
old_key = name.format(idx=old_idx)
new_key = name.format(idx=i)
new_state_dict[new_key] = state_dict[old_key].clone()
# Then, merge information from removed blocks
for i in range(28):
if i not in keep_blocks:
# Find the nearest kept blocks
prev_kept = max([b for b in keep_blocks if b < i])
next_kept = min([b for b in keep_blocks if b > i])
# Calculate the weight based on position
weight = (i - prev_kept) / (next_kept - prev_kept)
for name in block_names:
removed_key = name.format(idx=i)
prev_new_key = name.format(idx=keep_blocks.index(prev_kept))
next_new_key = name.format(idx=keep_blocks.index(next_kept))
# Weighted merge for previous kept block
new_state_dict[prev_new_key] = weighted_merge(new_state_dict[prev_new_key], state_dict[removed_key], weight)
# Weighted merge for next kept block
new_state_dict[next_new_key] = weighted_merge(new_state_dict[next_new_key], state_dict[removed_key],
1 - weight)
# Convert to fp16 and move to CPU
for key, value in new_state_dict.items():
new_state_dict[key] = value.to(torch.float16).cpu()
# Save the new state dict
save_file(new_state_dict, output_path, metadata=meta)
new_param_count = sum([v.numel() for v in new_state_dict.values()])
old_param_count = sum([v.numel() for v in state_dict.values()])
print(f"Old param count: {old_param_count:,}")
print(f"New param count: {new_param_count:,}")

View File

@@ -0,0 +1,84 @@
import torch
from safetensors.torch import load_file, save_file
from collections import OrderedDict
meta = OrderedDict()
meta['format'] = "pt"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def reduce_weight(weight, target_size):
weight = weight.to(device, torch.float32)
original_shape = weight.shape
flattened = weight.view(-1, original_shape[-1])
if flattened.shape[1] <= target_size:
return weight
U, S, V = torch.svd(flattened)
reduced = torch.mm(U[:, :target_size], torch.diag(S[:target_size]))
if reduced.shape[1] < target_size:
padding = torch.zeros(reduced.shape[0], target_size - reduced.shape[1], device=device)
reduced = torch.cat((reduced, padding), dim=1)
return reduced.view(original_shape[:-1] + (target_size,))
def reduce_bias(bias, target_size):
bias = bias.to(device, torch.float32)
original_size = bias.shape[0]
if original_size <= target_size:
return torch.nn.functional.pad(bias, (0, target_size - original_size))
else:
return bias.view(-1, original_size // target_size).mean(dim=1)[:target_size]
# Load your original state dict
state_dict = load_file(
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
# Create a new state dict for the reduced model
new_state_dict = {}
source_hidden_size = 1152
target_hidden_size = 1024
for key, value in state_dict.items():
value = value.to(device, torch.float32)
if 'weight' in key or 'scale_shift_table' in key:
if value.shape[0] == source_hidden_size:
value = value[:target_hidden_size]
elif value.shape[0] == source_hidden_size * 4:
value = value[:target_hidden_size * 4]
elif value.shape[0] == source_hidden_size * 6:
value = value[:target_hidden_size * 6]
if len(value.shape) > 1 and value.shape[
1] == source_hidden_size and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
value = value[:, :target_hidden_size]
elif len(value.shape) > 1 and value.shape[1] == source_hidden_size * 4:
value = value[:, :target_hidden_size * 4]
elif 'bias' in key:
if value.shape[0] == source_hidden_size:
value = value[:target_hidden_size]
elif value.shape[0] == source_hidden_size * 4:
value = value[:target_hidden_size * 4]
elif value.shape[0] == source_hidden_size * 6:
value = value[:target_hidden_size * 6]
new_state_dict[key] = value
# Move all to CPU and convert to float16
for key, value in new_state_dict.items():
new_state_dict[key] = value.cpu().to(torch.float16)
# Save the new state dict
save_file(new_state_dict,
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
metadata=meta)
print("Done!")

View File

@@ -0,0 +1,110 @@
import torch
from safetensors.torch import load_file, save_file
from collections import OrderedDict
meta = OrderedDict()
meta['format'] = "pt"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def reduce_weight(weight, target_size):
weight = weight.to(device, torch.float32)
original_shape = weight.shape
if len(original_shape) == 1:
# For 1D tensors, simply truncate
return weight[:target_size]
if original_shape[0] <= target_size:
return weight
# Reshape the tensor to 2D
flattened = weight.reshape(original_shape[0], -1)
# Perform SVD
U, S, V = torch.svd(flattened)
# Reduce the dimensions
reduced = torch.mm(U[:target_size, :], torch.diag(S)).mm(V.t())
# Reshape back to the original shape with reduced first dimension
new_shape = (target_size,) + original_shape[1:]
return reduced.reshape(new_shape)
def reduce_bias(bias, target_size):
bias = bias.to(device, torch.float32)
return bias[:target_size]
# Load your original state dict
state_dict = load_file(
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
# Create a new state dict for the reduced model
new_state_dict = {}
for key, value in state_dict.items():
value = value.to(device, torch.float32)
if 'weight' in key or 'scale_shift_table' in key:
if value.shape[0] == 1152:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
# reshape to (1152, -1)
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 512)
value = value.view(output_shape)
else:
# value = reduce_weight(value.t(), 576).t().contiguous()
value = reduce_weight(value, 512)
pass
elif value.shape[0] == 4608:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 2048)
value = value.view(output_shape)
else:
value = reduce_weight(value, 2048)
elif value.shape[0] == 6912:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 3072)
value = value.view(output_shape)
else:
value = reduce_weight(value, 3072)
if len(value.shape) > 1 and value.shape[
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
pass
elif len(value.shape) > 1 and value.shape[1] == 4608:
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
pass
elif 'bias' in key:
if value.shape[0] == 1152:
value = reduce_bias(value, 512)
elif value.shape[0] == 4608:
value = reduce_bias(value, 2048)
elif value.shape[0] == 6912:
value = reduce_bias(value, 3072)
new_state_dict[key] = value
# Move all to CPU and convert to float16
for key, value in new_state_dict.items():
new_state_dict[key] = value.cpu().to(torch.float16)
# Save the new state dict
save_file(new_state_dict,
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
metadata=meta)
print("Done!")

View File

@@ -0,0 +1,100 @@
import torch
from safetensors.torch import load_file, save_file
from collections import OrderedDict
meta = OrderedDict()
meta['format'] = "pt"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def reduce_weight(weight, target_size):
weight = weight.to(device, torch.float32)
# resize so target_size is the first dimension
tmp_weight = weight.view(1, 1, weight.shape[0], weight.shape[1])
# use interpolate to resize the tensor
new_weight = torch.nn.functional.interpolate(tmp_weight, size=(target_size, weight.shape[1]), mode='bicubic', align_corners=True)
# reshape back to original shape
return new_weight.view(target_size, weight.shape[1])
def reduce_bias(bias, target_size):
bias = bias.view(1, 1, bias.shape[0], 1)
new_bias = torch.nn.functional.interpolate(bias, size=(target_size, 1), mode='bicubic', align_corners=True)
return new_bias.view(target_size)
# Load your original state dict
state_dict = load_file(
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.orig.safetensors")
# Create a new state dict for the reduced model
new_state_dict = {}
for key, value in state_dict.items():
value = value.to(device, torch.float32)
if 'weight' in key or 'scale_shift_table' in key:
if value.shape[0] == 1152:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (512, orig_shape[1], orig_shape[2], orig_shape[3]) # reshape to (1152, -1)
# reshape to (1152, -1)
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 512)
value = value.view(output_shape)
else:
# value = reduce_weight(value.t(), 576).t().contiguous()
value = reduce_weight(value, 512)
pass
elif value.shape[0] == 4608:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (2048, orig_shape[1], orig_shape[2], orig_shape[3])
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 2048)
value = value.view(output_shape)
else:
value = reduce_weight(value, 2048)
elif value.shape[0] == 6912:
if len(value.shape) == 4:
orig_shape = value.shape
output_shape = (3072, orig_shape[1], orig_shape[2], orig_shape[3])
value = value.view(value.shape[0], -1)
value = reduce_weight(value, 3072)
value = value.view(output_shape)
else:
value = reduce_weight(value, 3072)
if len(value.shape) > 1 and value.shape[
1] == 1152 and 'attn2.to_k.weight' not in key and 'attn2.to_v.weight' not in key:
value = reduce_weight(value.t(), 512).t().contiguous() # Transpose before and after reduction
pass
elif len(value.shape) > 1 and value.shape[1] == 4608:
value = reduce_weight(value.t(), 2048).t().contiguous() # Transpose before and after reduction
pass
elif 'bias' in key:
if value.shape[0] == 1152:
value = reduce_bias(value, 512)
elif value.shape[0] == 4608:
value = reduce_bias(value, 2048)
elif value.shape[0] == 6912:
value = reduce_bias(value, 3072)
new_state_dict[key] = value
# Move all to CPU and convert to float16
for key, value in new_state_dict.items():
new_state_dict[key] = value.cpu().to(torch.float16)
# Save the new state dict
save_file(new_state_dict,
"/home/jaret/Dev/models/hf/PixArt-Sigma-XL-2-512_MS_t5large_raw/transformer/diffusion_pytorch_model.safetensors",
metadata=meta)
print("Done!")

View File

@@ -7,11 +7,13 @@ from torchvision import transforms
import sys
import os
import cv2
import random
from transformers import CLIPImageProcessor
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from toolkit.paths import SD_SCRIPTS_ROOT
from toolkit.image_utils import show_img
import torchvision.transforms.functional
from toolkit.image_utils import save_tensors, show_img, show_tensors
sys.path.append(SD_SCRIPTS_ROOT)
@@ -21,83 +23,118 @@ from toolkit.data_loader import AiToolkitDataset, get_dataloader_from_datasets,
trigger_dataloader_setup_epoch
from toolkit.config_modules import DatasetConfig
import argparse
from tqdm import tqdm
parser = argparse.ArgumentParser()
parser.add_argument('dataset_folder', type=str, default='input')
parser.add_argument('--epochs', type=int, default=1)
parser.add_argument('--num_frames', type=int, default=1)
parser.add_argument('--output_path', type=str, default=None)
args = parser.parse_args()
if args.output_path is not None:
args.output_path = os.path.abspath(args.output_path)
os.makedirs(args.output_path, exist_ok=True)
dataset_folder = args.dataset_folder
resolution = 1024
resolution = 512
bucket_tolerance = 64
batch_size = 1
clip_processor = CLIPImageProcessor.from_pretrained("openai/clip-vit-base-patch16")
class FakeAdapter:
def __init__(self):
self.clip_image_processor = clip_processor
## make fake sd
class FakeSD:
def __init__(self):
self.adapter = FakeAdapter()
##
dataset_config = DatasetConfig(
dataset_path=dataset_folder,
# clip_image_path=dataset_folder,
# square_crop=True,
resolution=resolution,
caption_ext='json',
# caption_ext='json',
default_caption='default',
# clip_image_path='/mnt/Datasets2/regs/yetibear_xl_v14/random_aspect/',
buckets=True,
bucket_tolerance=bucket_tolerance,
poi='person',
augmentations=[
{
'method': 'RandomBrightnessContrast',
'brightness_limit': (-0.3, 0.3),
'contrast_limit': (-0.3, 0.3),
'brightness_by_max': False,
'p': 1.0
},
{
'method': 'HueSaturationValue',
'hue_shift_limit': (-0, 0),
'sat_shift_limit': (-40, 40),
'val_shift_limit': (-40, 40),
'p': 1.0
},
# {
# 'method': 'RGBShift',
# 'r_shift_limit': (-20, 20),
# 'g_shift_limit': (-20, 20),
# 'b_shift_limit': (-20, 20),
# 'p': 1.0
# },
]
shrink_video_to_frames=True,
num_frames=args.num_frames,
# poi='person',
# shuffle_augmentations=True,
# augmentations=[
# {
# 'method': 'Posterize',
# 'num_bits': [(0, 4), (0, 4), (0, 4)],
# 'p': 1.0
# },
#
# ]
)
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size)
dataloader: DataLoader = get_dataloader_from_datasets([dataset_config], batch_size=batch_size, sd=FakeSD())
# run through an epoch ang check sizes
dataloader_iterator = iter(dataloader)
idx = 0
for epoch in range(args.epochs):
for batch in dataloader:
for batch in tqdm(dataloader):
batch: 'DataLoaderBatchDTO'
img_batch = batch.tensor
frames = 1
if len(img_batch.shape) == 5:
frames = img_batch.shape[1]
batch_size, frames, channels, height, width = img_batch.shape
else:
batch_size, channels, height, width = img_batch.shape
chunks = torch.chunk(img_batch, batch_size, dim=0)
# put them so they are size by side
big_img = torch.cat(chunks, dim=3)
big_img = big_img.squeeze(0)
# img_batch = color_block_imgs(img_batch, neg1_1=True)
min_val = big_img.min()
max_val = big_img.max()
# chunks = torch.chunk(img_batch, batch_size, dim=0)
# # put them so they are size by side
# big_img = torch.cat(chunks, dim=3)
# big_img = big_img.squeeze(0)
#
# control_chunks = torch.chunk(batch.clip_image_tensor, batch_size, dim=0)
# big_control_img = torch.cat(control_chunks, dim=3)
# big_control_img = big_control_img.squeeze(0) * 2 - 1
#
#
# # resize control image
# big_control_img = torchvision.transforms.Resize((width, height))(big_control_img)
#
# big_img = torch.cat([big_img, big_control_img], dim=2)
#
# min_val = big_img.min()
# max_val = big_img.max()
#
# big_img = (big_img / 2 + 0.5).clamp(0, 1)
big_img = (big_img / 2 + 0.5).clamp(0, 1)
big_img = img_batch
# big_img = big_img.clamp(-1, 1)
if args.output_path is not None:
save_tensors(big_img, os.path.join(args.output_path, f'{idx}.png'))
else:
show_tensors(big_img)
# convert to image
img = transforms.ToPILImage()(big_img)
# convert to image
# img = transforms.ToPILImage()(big_img)
#
# show_img(img)
show_img(img)
time.sleep(1.0)
time.sleep(0.2)
idx += 1
# if not last epoch
if epoch < args.epochs - 1:
trigger_dataloader_setup_epoch(dataloader)

130
testing/test_vae.py Normal file
View File

@@ -0,0 +1,130 @@
import argparse
import os
from PIL import Image
import torch
from torchvision.transforms import Resize, ToTensor
from diffusers import AutoencoderKL
from pytorch_fid import fid_score
from skimage.metrics import peak_signal_noise_ratio as psnr
import lpips
from tqdm import tqdm
from torchvision import transforms
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
def load_images(folder_path):
images = []
for filename in os.listdir(folder_path):
if filename.lower().endswith(('.png', '.jpg', '.jpeg')):
img_path = os.path.join(folder_path, filename)
images.append(img_path)
return images
def paramiter_count(model):
state_dict = model.state_dict()
paramiter_count = 0
for key in state_dict:
paramiter_count += torch.numel(state_dict[key])
return int(paramiter_count)
def calculate_metrics(vae, images, max_imgs=-1, save_output=False):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
vae = vae.to(device)
lpips_model = lpips.LPIPS(net='alex').to(device)
rfid_scores = []
psnr_scores = []
lpips_scores = []
# transform = transforms.Compose([
# transforms.Resize(256, antialias=True),
# transforms.CenterCrop(256)
# ])
# needs values between -1 and 1
to_tensor = ToTensor()
# remove _reconstructed.png files
images = [img for img in images if not img.endswith("_reconstructed.png")]
if max_imgs > 0 and len(images) > max_imgs:
images = images[:max_imgs]
for img_path in tqdm(images):
try:
img = Image.open(img_path).convert('RGB')
# img_tensor = to_tensor(transform(img)).unsqueeze(0).to(device)
img_tensor = to_tensor(img).unsqueeze(0).to(device)
img_tensor = 2 * img_tensor - 1
# if width or height is not divisible by 8, crop it
if img_tensor.shape[2] % 8 != 0 or img_tensor.shape[3] % 8 != 0:
img_tensor = img_tensor[:, :, :img_tensor.shape[2] // 8 * 8, :img_tensor.shape[3] // 8 * 8]
except Exception as e:
print(f"Error processing {img_path}: {e}")
continue
with torch.no_grad():
reconstructed = vae.decode(vae.encode(img_tensor).latent_dist.sample()).sample
# Calculate rFID
# rfid = fid_score.calculate_frechet_distance(vae, img_tensor, reconstructed)
# rfid_scores.append(rfid)
# Calculate PSNR
psnr_val = psnr(img_tensor.cpu().numpy(), reconstructed.cpu().numpy())
psnr_scores.append(psnr_val)
# Calculate LPIPS
lpips_val = lpips_model(img_tensor, reconstructed).item()
lpips_scores.append(lpips_val)
# avg_rfid = sum(rfid_scores) / len(rfid_scores)
avg_rfid = 0
avg_psnr = sum(psnr_scores) / len(psnr_scores)
avg_lpips = sum(lpips_scores) / len(lpips_scores)
if save_output:
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
folder = os.path.dirname(img_path)
save_path = os.path.join(folder, filename_no_ext + "_reconstructed.png")
reconstructed = (reconstructed + 1) / 2
reconstructed = reconstructed.clamp(0, 1)
reconstructed = transforms.ToPILImage()(reconstructed[0].cpu())
reconstructed.save(save_path)
return avg_rfid, avg_psnr, avg_lpips
def main():
parser = argparse.ArgumentParser(description="Calculate average rFID, PSNR, and LPIPS for VAE reconstructions")
parser.add_argument("--vae_path", type=str, required=True, help="Path to the VAE model")
parser.add_argument("--image_folder", type=str, required=True, help="Path to the folder containing images")
parser.add_argument("--max_imgs", type=int, default=-1, help="Max num of images. Default is -1 for all images.")
# boolean store true
parser.add_argument("--save_output", action="store_true", help="Save the output images")
args = parser.parse_args()
if os.path.isfile(args.vae_path):
vae = AutoencoderKL.from_single_file(args.vae_path)
else:
try:
vae = AutoencoderKL.from_pretrained(args.vae_path)
except:
vae = AutoencoderKL.from_pretrained(args.vae_path, subfolder="vae")
vae.eval()
vae = vae.to(device)
print(f"Model has {paramiter_count(vae)} parameters")
images = load_images(args.image_folder)
avg_rfid, avg_psnr, avg_lpips = calculate_metrics(vae, images, args.max_imgs, args.save_output)
# print(f"Average rFID: {avg_rfid}")
print(f"Average PSNR: {avg_psnr}")
print(f"Average LPIPS: {avg_lpips}")
if __name__ == "__main__":
main()

17
toolkit/accelerator.py Normal file
View File

@@ -0,0 +1,17 @@
from accelerate import Accelerator
from diffusers.utils.torch_utils import is_compiled_module
global_accelerator = None
def get_accelerator() -> Accelerator:
global global_accelerator
if global_accelerator is None:
global_accelerator = Accelerator()
return global_accelerator
def unwrap_model(model):
accelerator = get_accelerator()
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model

55
toolkit/assistant_lora.py Normal file
View File

@@ -0,0 +1,55 @@
from typing import TYPE_CHECKING
from toolkit.config_modules import NetworkConfig
from toolkit.lora_special import LoRASpecialNetwork
from safetensors.torch import load_file
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
def load_assistant_lora_from_path(adapter_path, sd: 'StableDiffusion') -> LoRASpecialNetwork:
if not sd.is_flux:
raise ValueError("Only Flux models can load assistant adapters currently.")
pipe = sd.pipeline
print(f"Loading assistant adapter from {adapter_path}")
adapter_name = adapter_path.split("/")[-1].split(".")[0]
lora_state_dict = load_file(adapter_path)
linear_dim = int(lora_state_dict['transformer.single_transformer_blocks.0.attn.to_k.lora_A.weight'].shape[0])
# linear_alpha = int(lora_state_dict['lora_transformer_single_transformer_blocks_0_attn_to_k.alpha'].item())
linear_alpha = linear_dim
transformer_only = 'transformer.proj_out.alpha' not in lora_state_dict
# get dim and scale
network_config = NetworkConfig(
linear=linear_dim,
linear_alpha=linear_alpha,
transformer_only=transformer_only,
)
network = LoRASpecialNetwork(
text_encoder=pipe.text_encoder,
unet=pipe.transformer,
lora_dim=network_config.linear,
multiplier=1.0,
alpha=network_config.linear_alpha,
train_unet=True,
train_text_encoder=False,
is_flux=True,
network_config=network_config,
network_type=network_config.type,
transformer_only=network_config.transformer_only,
is_assistant_adapter=True
)
network.apply_to(
pipe.text_encoder,
pipe.transformer,
apply_text_encoder=False,
apply_unet=True
)
network.force_to(sd.device_torch, dtype=sd.torch_dtype)
network.eval()
network._update_torch_multiplier()
network.load_weights(lora_state_dict)
network.is_active = True
return network

View File

@@ -31,12 +31,18 @@ def get_mean_std(tensor):
def adain(content_features, style_features):
# Assumes that the content and style features are of shape (batch_size, channels, width, height)
dims = [2, 3]
if len(content_features.shape) == 3:
# content_features = content_features.unsqueeze(0)
# style_features = style_features.unsqueeze(0)
dims = [1]
# Step 1: Calculate mean and variance of content features
content_mean, content_var = torch.mean(content_features, dim=[2, 3], keepdim=True), torch.var(content_features,
dim=[2, 3],
content_mean, content_var = torch.mean(content_features, dim=dims, keepdim=True), torch.var(content_features,
dim=dims,
keepdim=True)
# Step 2: Calculate mean and variance of style features
style_mean, style_var = torch.mean(style_features, dim=[2, 3], keepdim=True), torch.var(style_features, dim=[2, 3],
style_mean, style_var = torch.mean(style_features, dim=dims, keepdim=True), torch.var(style_features, dim=dims,
keepdim=True)
# Step 3: Normalize content features

View File

@@ -51,9 +51,11 @@ resolutions_1024: List[BucketResolution] = [
{"width": 512, "height": 1920},
{"width": 512, "height": 1984},
{"width": 512, "height": 2048},
# extra wides
{"width": 8192, "height": 128},
{"width": 128, "height": 8192},
]
def get_bucket_sizes(resolution: int = 512, divisibility: int = 8) -> List[BucketResolution]:
# determine scaler form 1024 to resolution
scaler = resolution / 1024
@@ -124,4 +126,4 @@ def get_bucket_for_image_size(
if closest_bucket is None:
raise ValueError("No suitable bucket found")
return closest_bucket
return closest_bucket

View File

@@ -0,0 +1,406 @@
from typing import TYPE_CHECKING, Mapping, Any
import torch
import weakref
from toolkit.config_modules import AdapterConfig
from toolkit.models.clip_fusion import ZipperBlock
from toolkit.models.zipper_resampler import ZipperModule
from toolkit.prompt_utils import PromptEmbeds
from toolkit.train_tools import get_torch_dtype
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
from transformers import (
CLIPImageProcessor,
CLIPVisionModelWithProjection,
CLIPVisionModel
)
from toolkit.resampler import Resampler
import torch.nn as nn
class Embedder(nn.Module):
def __init__(
self,
num_input_tokens: int = 1,
input_dim: int = 1024,
num_output_tokens: int = 8,
output_dim: int = 768,
mid_dim: int = 1024
):
super(Embedder, self).__init__()
self.num_output_tokens = num_output_tokens
self.num_input_tokens = num_input_tokens
self.input_dim = input_dim
self.output_dim = output_dim
self.layer_norm = nn.LayerNorm(input_dim)
self.fc1 = nn.Linear(input_dim, mid_dim)
self.gelu = nn.GELU()
# self.fc2 = nn.Linear(mid_dim, mid_dim)
self.fc2 = nn.Linear(mid_dim, mid_dim)
self.fc2.weight.data.zero_()
self.layer_norm2 = nn.LayerNorm(mid_dim)
self.fc3 = nn.Linear(mid_dim, mid_dim)
self.gelu2 = nn.GELU()
self.fc4 = nn.Linear(mid_dim, output_dim * num_output_tokens)
# set the weights to 0
self.fc3.weight.data.zero_()
self.fc4.weight.data.zero_()
# self.static_tokens = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
# self.scaler = nn.Parameter(torch.zeros(num_output_tokens, output_dim))
def forward(self, x):
if len(x.shape) == 2:
x = x.unsqueeze(1)
x = self.layer_norm(x)
x = self.fc1(x)
x = self.gelu(x)
x = self.fc2(x)
x = self.layer_norm2(x)
x = self.fc3(x)
x = self.gelu2(x)
x = self.fc4(x)
x = x.view(-1, self.num_output_tokens, self.output_dim)
return x
class ClipVisionAdapter(torch.nn.Module):
def __init__(self, sd: 'StableDiffusion', adapter_config: AdapterConfig):
super().__init__()
self.config = adapter_config
self.trigger = adapter_config.trigger
self.trigger_class_name = adapter_config.trigger_class_name
self.sd_ref: weakref.ref = weakref.ref(sd)
# embedding stuff
self.text_encoder_list = sd.text_encoder if isinstance(sd.text_encoder, list) else [sd.text_encoder]
self.tokenizer_list = sd.tokenizer if isinstance(sd.tokenizer, list) else [sd.tokenizer]
placeholder_tokens = [self.trigger]
# add dummy tokens for multi-vector
additional_tokens = []
for i in range(1, self.config.num_tokens):
additional_tokens.append(f"{self.trigger}_{i}")
placeholder_tokens += additional_tokens
# handle dual tokenizer
self.tokenizer_list = self.sd_ref().tokenizer if isinstance(self.sd_ref().tokenizer, list) else [
self.sd_ref().tokenizer]
self.text_encoder_list = self.sd_ref().text_encoder if isinstance(self.sd_ref().text_encoder, list) else [
self.sd_ref().text_encoder]
self.placeholder_token_ids = []
self.embedding_tokens = []
print(f"Adding {placeholder_tokens} tokens to tokenizer")
print(f"Adding {self.config.num_tokens} tokens to tokenizer")
for text_encoder, tokenizer in zip(self.text_encoder_list, self.tokenizer_list):
num_added_tokens = tokenizer.add_tokens(placeholder_tokens)
if num_added_tokens != self.config.num_tokens:
raise ValueError(
f"The tokenizer already contains the token {self.trigger}. Please pass a different"
f" `placeholder_token` that is not already in the tokenizer. Only added {num_added_tokens}"
)
# Convert the initializer_token, placeholder_token to ids
init_token_ids = tokenizer.encode(self.config.trigger_class_name, add_special_tokens=False)
# if length of token ids is more than number of orm embedding tokens fill with *
if len(init_token_ids) > self.config.num_tokens:
init_token_ids = init_token_ids[:self.config.num_tokens]
elif len(init_token_ids) < self.config.num_tokens:
pad_token_id = tokenizer.encode(["*"], add_special_tokens=False)
init_token_ids += pad_token_id * (self.config.num_tokens - len(init_token_ids))
placeholder_token_ids = tokenizer.encode(placeholder_tokens, add_special_tokens=False)
self.placeholder_token_ids.append(placeholder_token_ids)
# Resize the token embeddings as we are adding new special tokens to the tokenizer
text_encoder.resize_token_embeddings(len(tokenizer))
# Initialise the newly added placeholder token with the embeddings of the initializer token
token_embeds = text_encoder.get_input_embeddings().weight.data
with torch.no_grad():
for initializer_token_id, token_id in zip(init_token_ids, placeholder_token_ids):
token_embeds[token_id] = token_embeds[initializer_token_id].clone()
# replace "[name] with this. on training. This is automatically generated in pipeline on inference
self.embedding_tokens.append(" ".join(tokenizer.convert_ids_to_tokens(placeholder_token_ids)))
# backup text encoder embeddings
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
try:
self.clip_image_processor = CLIPImageProcessor.from_pretrained(self.config.image_encoder_path)
except EnvironmentError:
self.clip_image_processor = CLIPImageProcessor()
self.device = self.sd_ref().unet.device
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
self.config.image_encoder_path,
ignore_mismatched_sizes=True
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
if self.config.train_image_encoder:
self.image_encoder.train()
else:
self.image_encoder.eval()
# max_seq_len = CLIP tokens + CLS token
image_encoder_state_dict = self.image_encoder.state_dict()
in_tokens = 257
if "vision_model.embeddings.position_embedding.weight" in image_encoder_state_dict:
# clip
in_tokens = int(image_encoder_state_dict["vision_model.embeddings.position_embedding.weight"].shape[0])
if hasattr(self.image_encoder.config, 'hidden_sizes'):
embedding_dim = self.image_encoder.config.hidden_sizes[-1]
else:
embedding_dim = self.image_encoder.config.target_hidden_size
if self.config.clip_layer == 'image_embeds':
in_tokens = 1
embedding_dim = self.image_encoder.config.projection_dim
self.embedder = Embedder(
num_output_tokens=self.config.num_tokens,
num_input_tokens=in_tokens,
input_dim=embedding_dim,
output_dim=self.sd_ref().unet.config['cross_attention_dim'],
mid_dim=embedding_dim * self.config.num_tokens,
).to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
self.embedder.train()
def state_dict(self, *args, destination=None, prefix='', keep_vars=False):
state_dict = {
'embedder': self.embedder.state_dict(*args, destination=destination, prefix=prefix, keep_vars=keep_vars)
}
if self.config.train_image_encoder:
state_dict['image_encoder'] = self.image_encoder.state_dict(
*args, destination=destination, prefix=prefix,
keep_vars=keep_vars)
return state_dict
def load_state_dict(self, state_dict: Mapping[str, Any], strict: bool = True):
self.embedder.load_state_dict(state_dict["embedder"], strict=strict)
if self.config.train_image_encoder and 'image_encoder' in state_dict:
self.image_encoder.load_state_dict(state_dict["image_encoder"], strict=strict)
def parameters(self, *args, **kwargs):
yield from self.embedder.parameters(*args, **kwargs)
def named_parameters(self, *args, **kwargs):
yield from self.embedder.named_parameters(*args, **kwargs)
def get_clip_image_embeds_from_tensors(
self, tensors_0_1: torch.Tensor, drop=False,
is_training=False,
has_been_preprocessed=False
) -> torch.Tensor:
with torch.no_grad():
if not has_been_preprocessed:
# tensors should be 0-1
if tensors_0_1.ndim == 3:
tensors_0_1 = tensors_0_1.unsqueeze(0)
# training tensors are 0 - 1
tensors_0_1 = tensors_0_1.to(self.device, dtype=torch.float16)
# if images are out of this range throw error
if tensors_0_1.min() < -0.3 or tensors_0_1.max() > 1.3:
raise ValueError("image tensor values must be between 0 and 1. Got min: {}, max: {}".format(
tensors_0_1.min(), tensors_0_1.max()
))
# unconditional
if drop:
if self.clip_noise_zero:
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
dtype=get_torch_dtype(self.sd_ref().dtype))
tensors_0_1 = tensors_0_1 * noise_scale
else:
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
# tensors_0_1 = tensors_0_1 * 0
clip_image = self.clip_image_processor(
images=tensors_0_1,
return_tensors="pt",
do_resize=True,
do_rescale=False,
).pixel_values
else:
if drop:
# scale the noise down
if self.clip_noise_zero:
tensors_0_1 = torch.rand_like(tensors_0_1).detach()
noise_scale = torch.rand([tensors_0_1.shape[0], 1, 1, 1], device=self.device,
dtype=get_torch_dtype(self.sd_ref().dtype))
tensors_0_1 = tensors_0_1 * noise_scale
else:
tensors_0_1 = torch.zeros_like(tensors_0_1).detach()
# tensors_0_1 = tensors_0_1 * 0
mean = torch.tensor(self.clip_image_processor.image_mean).to(
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
).detach()
std = torch.tensor(self.clip_image_processor.image_std).to(
self.device, dtype=get_torch_dtype(self.sd_ref().dtype)
).detach()
tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
clip_image = (tensors_0_1 - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
else:
clip_image = tensors_0_1
clip_image = clip_image.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype)).detach()
with torch.set_grad_enabled(is_training):
if is_training:
self.image_encoder.train()
else:
self.image_encoder.eval()
clip_output = self.image_encoder(clip_image, output_hidden_states=True)
if self.config.clip_layer == 'penultimate_hidden_states':
# they skip last layer for ip+
# https://github.com/tencent-ailab/IP-Adapter/blob/f4b6742db35ea6d81c7b829a55b0a312c7f5a677/tutorial_train_plus.py#L403C26-L403C26
clip_image_embeds = clip_output.hidden_states[-2]
elif self.config.clip_layer == 'last_hidden_state':
clip_image_embeds = clip_output.hidden_states[-1]
else:
clip_image_embeds = clip_output.image_embeds
return clip_image_embeds
import torch
def set_vec(self, new_vector, text_encoder_idx=0):
# Get the embedding layer
embedding_layer = self.text_encoder_list[text_encoder_idx].get_input_embeddings()
# Indices to replace in the embeddings
indices_to_replace = self.placeholder_token_ids[text_encoder_idx]
# Replace the specified embeddings with new_vector
for idx in indices_to_replace:
vector_idx = idx - indices_to_replace[0]
embedding_layer.weight[idx] = new_vector[vector_idx]
# adds it to the tokenizer
def forward(self, clip_image_embeds: torch.Tensor) -> PromptEmbeds:
clip_image_embeds = clip_image_embeds.to(self.device, dtype=get_torch_dtype(self.sd_ref().dtype))
if clip_image_embeds.ndim == 2:
# expand the token dimension
clip_image_embeds = clip_image_embeds.unsqueeze(1)
image_prompt_embeds = self.embedder(clip_image_embeds)
# todo add support for multiple batch sizes
if image_prompt_embeds.shape[0] != 1:
raise ValueError("Batch size must be 1 for embedder for now")
# output on sd1.5 is bs, num_tokens, 768
if len(self.text_encoder_list) == 1:
# add it to the text encoder
self.set_vec(image_prompt_embeds[0], text_encoder_idx=0)
elif len(self.text_encoder_list) == 2:
if self.text_encoder_list[0].config.target_hidden_size + self.text_encoder_list[1].config.target_hidden_size != \
image_prompt_embeds.shape[2]:
raise ValueError("Something went wrong. The embeddings do not match the text encoder sizes")
# sdxl variants
# image_prompt_embeds = 2048
# te1 = 768
# te2 = 1280
te1_embeds = image_prompt_embeds[:, :, :self.text_encoder_list[0].config.target_hidden_size]
te2_embeds = image_prompt_embeds[:, :, self.text_encoder_list[0].config.target_hidden_size:]
self.set_vec(te1_embeds[0], text_encoder_idx=0)
self.set_vec(te2_embeds[0], text_encoder_idx=1)
else:
raise ValueError("Unsupported number of text encoders")
# just a place to put a breakpoint
pass
def restore_embeddings(self):
# Let's make sure we don't update any embedding weights besides the newly added token
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(
self.text_encoder_list,
self.tokenizer_list,
self.orig_embeds_params,
self.placeholder_token_ids
):
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
index_no_updates[
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
with torch.no_grad():
text_encoder.get_input_embeddings().weight[
index_no_updates
] = orig_embeds[index_no_updates]
# detach it all
text_encoder.get_input_embeddings().weight.detach_()
def enable_gradient_checkpointing(self):
self.image_encoder.gradient_checkpointing = True
def inject_trigger_into_prompt(self, prompt, expand_token=False, to_replace_list=None, add_if_not_present=True):
output_prompt = prompt
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
default_replacements = ["[name]", "[trigger]"]
replace_with = embedding_tokens if expand_token else self.trigger
if to_replace_list is None:
to_replace_list = default_replacements
else:
to_replace_list += default_replacements
# remove duplicates
to_replace_list = list(set(to_replace_list))
# replace them all
for to_replace in to_replace_list:
# replace it
output_prompt = output_prompt.replace(to_replace, replace_with)
# see how many times replace_with is in the prompt
num_instances = output_prompt.count(replace_with)
if num_instances == 0 and add_if_not_present:
# add it to the beginning of the prompt
output_prompt = replace_with + " " + output_prompt
if num_instances > 1:
print(
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
return output_prompt
# reverses injection with class name. useful for normalizations
def inject_trigger_class_name_into_prompt(self, prompt):
output_prompt = prompt
embedding_tokens = self.embedding_tokens[0] # shoudl be the same
default_replacements = ["[name]", "[trigger]", embedding_tokens, self.trigger]
replace_with = self.config.trigger_class_name
to_replace_list = default_replacements
# remove duplicates
to_replace_list = list(set(to_replace_list))
# replace them all
for to_replace in to_replace_list:
# replace it
output_prompt = output_prompt.replace(to_replace, replace_with)
# see how many times replace_with is in the prompt
num_instances = output_prompt.count(replace_with)
if num_instances > 1:
print(
f"Warning: {replace_with} token appears {num_instances} times in prompt {output_prompt}. This may cause issues.")
return output_prompt

View File

@@ -43,9 +43,7 @@ def preprocess_config(config: OrderedDict, name: str = None):
if "name" not in config["config"] and name is None:
raise ValueError("config file must have a config.name key")
# we need to replace tags. For now just [name]
if name is not None:
config["config"]["name"] = name
else:
if name is None:
name = config["config"]["name"]
config_string = json.dumps(config)
config_string = config_string.replace("[name]", name)

View File

@@ -1,6 +1,6 @@
import os
import time
from typing import List, Optional, Literal, Union
from typing import List, Optional, Literal, Union, TYPE_CHECKING, Dict
import random
import torch
@@ -11,22 +11,31 @@ ImgExt = Literal['jpg', 'png', 'webp']
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):
self.save_every: int = kwargs.get('save_every', 1000)
self.dtype: str = kwargs.get('save_dtype', 'float16')
self.dtype: str = kwargs.get('dtype', 'float16')
self.max_step_saves_to_keep: int = kwargs.get('max_step_saves_to_keep', 5)
self.save_format: SaveFormat = kwargs.get('save_format', 'safetensors')
if self.save_format not in ['safetensors', 'diffusers']:
raise ValueError(f"save_format must be safetensors or diffusers, got {self.save_format}")
self.push_to_hub: bool = kwargs.get("push_to_hub", False)
self.hf_repo_id: Optional[str] = kwargs.get("hf_repo_id", None)
self.hf_private: Optional[str] = kwargs.get("hf_private", False)
class LogingConfig:
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:
@@ -45,7 +54,14 @@ class SampleConfig:
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
self.ext: ImgExt = kwargs.get('format', 'jpg')
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5) # step to start using refiner on sample if it exists
self.refiner_start_at = kwargs.get('refiner_start_at',
0.5) # step to start using refiner on sample if it exists
self.extra_values = kwargs.get('extra_values', [])
self.num_frames = kwargs.get('num_frames', 1)
self.fps: int = kwargs.get('fps', 16)
if self.num_frames > 1 and self.ext not in ['webp']:
print("Changing sample extention to animated webp")
self.ext = 'webp'
class LormModuleSettingsConfig:
@@ -90,7 +106,7 @@ class LoRMConfig:
})
NetworkType = Literal['lora', 'locon', 'lorm']
NetworkType = Literal['lora', 'locon', 'lorm', 'lokr']
class NetworkConfig:
@@ -109,6 +125,7 @@ class NetworkConfig:
self.linear_alpha: float = kwargs.get('linear_alpha', self.alpha)
self.conv_alpha: float = kwargs.get('conv_alpha', self.conv)
self.dropout: Union[float, None] = kwargs.get('dropout', None)
self.network_kwargs: dict = kwargs.get('network_kwargs', {})
self.lorm_config: Union[LoRMConfig, None] = None
lorm = kwargs.get('lorm', None)
@@ -122,24 +139,120 @@ class NetworkConfig:
if self.lorm_config.do_conv:
self.conv = 4
self.transformer_only = kwargs.get('transformer_only', True)
self.lokr_full_rank = kwargs.get('lokr_full_rank', False)
if self.lokr_full_rank and self.type.lower() == 'lokr':
self.linear = 9999999999
self.linear_alpha = 9999999999
self.conv = 9999999999
self.conv_alpha = 9999999999
# -1 automatically finds the largest factor
self.lokr_factor = kwargs.get('lokr_factor', -1)
AdapterTypes = Literal['t2i', 'ip', 'ip+']
AdapterTypes = Literal['t2i', 'ip', 'ip+', 'clip', 'ilora', 'photo_maker', 'control_net', 'control_lora']
CLIPLayer = Literal['penultimate_hidden_states', 'image_embeds', 'last_hidden_state']
class AdapterConfig:
def __init__(self, **kwargs):
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip
self.type: AdapterTypes = kwargs.get('type', 't2i') # t2i, ip, clip, control_net
self.in_channels: int = kwargs.get('in_channels', 3)
self.channels: List[int] = kwargs.get('channels', [320, 640, 1280, 1280])
self.num_res_blocks: int = kwargs.get('num_res_blocks', 2)
self.downscale_factor: int = kwargs.get('downscale_factor', 8)
self.adapter_type: str = kwargs.get('adapter_type', 'full_adapter')
self.image_dir: str = kwargs.get('image_dir', None)
self.test_img_path: str = kwargs.get('test_img_path', None)
self.test_img_path: List[str] = kwargs.get('test_img_path', None)
if self.test_img_path is not None:
if isinstance(self.test_img_path, str):
self.test_img_path = self.test_img_path.split(',')
self.test_img_path = [p.strip() for p in self.test_img_path]
self.test_img_path = [p for p in self.test_img_path if p != '']
self.train: str = kwargs.get('train', False)
self.image_encoder_path: str = kwargs.get('image_encoder_path', None)
self.name_or_path = kwargs.get('name_or_path', None)
num_tokens = kwargs.get('num_tokens', None)
if num_tokens is None and self.type.startswith('ip'):
if self.type == 'ip+':
num_tokens = 16
num_tokens = 16
elif self.type == 'ip':
num_tokens = 4
self.num_tokens: int = num_tokens
self.train_image_encoder: bool = kwargs.get('train_image_encoder', False)
self.train_only_image_encoder: bool = kwargs.get('train_only_image_encoder', False)
if self.train_only_image_encoder:
self.train_image_encoder = True
self.train_only_image_encoder_positional_embedding: bool = kwargs.get(
'train_only_image_encoder_positional_embedding', False)
self.image_encoder_arch: str = kwargs.get('image_encoder_arch', 'clip') # clip vit vit_hybrid, safe
self.safe_reducer_channels: int = kwargs.get('safe_reducer_channels', 512)
self.safe_channels: int = kwargs.get('safe_channels', 2048)
self.safe_tokens: int = kwargs.get('safe_tokens', 8)
self.quad_image: bool = kwargs.get('quad_image', False)
# clip vision
self.trigger = kwargs.get('trigger', 'tri993r')
self.trigger_class_name = kwargs.get('trigger_class_name', None)
self.class_names = kwargs.get('class_names', [])
self.clip_layer: CLIPLayer = kwargs.get('clip_layer', None)
if self.clip_layer is None:
if self.type.startswith('ip+'):
self.clip_layer = 'penultimate_hidden_states'
else:
self.clip_layer = 'last_hidden_state'
# text encoder
self.text_encoder_path: str = kwargs.get('text_encoder_path', None)
self.text_encoder_arch: str = kwargs.get('text_encoder_arch', 'clip') # clip t5
self.train_scaler: bool = kwargs.get('train_scaler', False)
self.scaler_lr: Optional[float] = kwargs.get('scaler_lr', None)
# trains with a scaler to easy channel bias but merges it in on save
self.merge_scaler: bool = kwargs.get('merge_scaler', False)
# for ilora
self.head_dim: int = kwargs.get('head_dim', 1024)
self.num_heads: int = kwargs.get('num_heads', 1)
self.ilora_down: bool = kwargs.get('ilora_down', True)
self.ilora_mid: bool = kwargs.get('ilora_mid', True)
self.ilora_up: bool = kwargs.get('ilora_up', True)
self.pixtral_max_image_size: int = kwargs.get('pixtral_max_image_size', 512)
self.pixtral_random_image_size: int = kwargs.get('pixtral_random_image_size', False)
self.flux_only_double: bool = kwargs.get('flux_only_double', False)
# train and use a conv layer to pool the embedding
self.conv_pooling: bool = kwargs.get('conv_pooling', False)
self.conv_pooling_stacks: int = kwargs.get('conv_pooling_stacks', 1)
self.sparse_autoencoder_dim: Optional[int] = kwargs.get('sparse_autoencoder_dim', None)
# for llm adapter
self.num_cloned_blocks: int = kwargs.get('num_cloned_blocks', 0)
self.quantize_llm: bool = kwargs.get('quantize_llm', False)
# for control lora only
lora_config: dict = kwargs.get('lora_config', None)
if lora_config is not None:
self.lora_config: NetworkConfig = NetworkConfig(**lora_config)
else:
self.lora_config = None
self.num_control_images: int = kwargs.get('num_control_images', 1)
# decimal for how often the control is dropped out and replaced with noise 1.0 is 100%
self.control_image_dropout: float = kwargs.get('control_image_dropout', 0.0)
self.has_inpainting_input: bool = kwargs.get('has_inpainting_input', False)
self.invert_inpaint_mask_chance: float = kwargs.get('invert_inpaint_mask_chance', 0.0)
class EmbeddingConfig:
def __init__(self, **kwargs):
@@ -147,6 +260,12 @@ class EmbeddingConfig:
self.tokens = kwargs.get('tokens', 4)
self.init_words = kwargs.get('init_words', '*')
self.save_format = kwargs.get('save_format', 'safetensors')
self.trigger_class_name = kwargs.get('trigger_class_name', None) # used for inverted masked prior
class DecoratorConfig:
def __init__(self, **kwargs):
self.num_tokens: str = kwargs.get('num_tokens', 4)
ContentOrStyleType = Literal['balanced', 'style', 'content']
@@ -157,6 +276,7 @@ class TrainConfig:
def __init__(self, **kwargs):
self.noise_scheduler = kwargs.get('noise_scheduler', 'ddpm')
self.content_or_style: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
self.content_or_style_reg: ContentOrStyleType = kwargs.get('content_or_style', 'balanced')
self.steps: int = kwargs.get('steps', 1000)
self.lr = kwargs.get('lr', 1e-6)
self.unet_lr = kwargs.get('unet_lr', self.lr)
@@ -171,12 +291,15 @@ 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)
self.train_unet = kwargs.get('train_unet', True)
self.train_text_encoder = kwargs.get('train_text_encoder', True)
self.train_text_encoder = kwargs.get('train_text_encoder', False)
self.train_refiner = kwargs.get('train_refiner', True)
self.train_turbo = kwargs.get('train_turbo', False)
self.show_turbo_outputs = kwargs.get('show_turbo_outputs', False)
self.min_snr_gamma = kwargs.get('min_snr_gamma', None)
self.snr_gamma = kwargs.get('snr_gamma', None)
# trains a gamma, offset, and scale to adjust loss to adapt to timestep differentials
@@ -184,6 +307,7 @@ class TrainConfig:
self.learnable_snr_gos = kwargs.get('learnable_snr_gos', False)
self.noise_offset = kwargs.get('noise_offset', 0.0)
self.skip_first_sample = kwargs.get('skip_first_sample', False)
self.force_first_sample = kwargs.get('force_first_sample', False)
self.gradient_checkpointing = kwargs.get('gradient_checkpointing', True)
self.weight_jitter = kwargs.get('weight_jitter', 0.0)
self.merge_network_on_save = kwargs.get('merge_network_on_save', False)
@@ -191,12 +315,20 @@ class TrainConfig:
self.start_step = kwargs.get('start_step', None)
self.free_u = kwargs.get('free_u', False)
self.adapter_assist_name_or_path: Optional[str] = kwargs.get('adapter_assist_name_or_path', None)
self.adapter_assist_type: Optional[str] = kwargs.get('adapter_assist_type', 't2i') # t2i, control_net
self.noise_multiplier = kwargs.get('noise_multiplier', 1.0)
self.target_noise_multiplier = kwargs.get('target_noise_multiplier', 1.0)
self.img_multiplier = kwargs.get('img_multiplier', 1.0)
self.noisy_latent_multiplier = kwargs.get('noisy_latent_multiplier', 1.0)
self.latent_multiplier = kwargs.get('latent_multiplier', 1.0)
self.negative_prompt = kwargs.get('negative_prompt', None)
self.max_negative_prompts = kwargs.get('max_negative_prompts', 1)
# multiplier applied to loos on regularization images
self.reg_weight = kwargs.get('reg_weight', 1.0)
self.num_train_timesteps = kwargs.get('num_train_timesteps', 1000)
self.random_noise_shift = kwargs.get('random_noise_shift', 0.0)
# automatically adapte the vae scaling based on the image norm
self.adaptive_scaling_factor = kwargs.get('adaptive_scaling_factor', False)
# dropout that happens before encoding. It functions independently per text encoder
self.prompt_dropout_prob = kwargs.get('prompt_dropout_prob', 0.0)
@@ -208,8 +340,16 @@ class TrainConfig:
# set to -1 to accumulate gradients for entire epoch
# warning, only do this with a small dataset or you will run out of memory
# This is legacy but left in for backwards compatibility
self.gradient_accumulation_steps = kwargs.get('gradient_accumulation_steps', 1)
# this will do proper gradient accumulation where you will not see a step until the end of the accumulation
# the method above will show a step every accumulation
self.gradient_accumulation = kwargs.get('gradient_accumulation', 1)
if self.gradient_accumulation > 1:
if self.gradient_accumulation_steps != 1:
raise ValueError("gradient_accumulation and gradient_accumulation_steps are mutually exclusive")
# short long captions will double your batch size. This only works when a dataset is
# prepared with a json caption file that has both short and long captions in it. It will
# Double up every image and run it through with both short and long captions. The idea
@@ -233,24 +373,123 @@ class TrainConfig:
# unmasked reign. It is unmasked regularization basically
self.inverted_mask_prior = kwargs.get('inverted_mask_prior', False)
self.inverted_mask_prior_multiplier = kwargs.get('inverted_mask_prior_multiplier', 0.5)
# DOP will will run the same image and prompt through the network without the trigger word blank and use it as a target
self.diff_output_preservation = kwargs.get('diff_output_preservation', False)
self.diff_output_preservation_multiplier = kwargs.get('diff_output_preservation_multiplier', 1.0)
# If the trigger word is in the prompt, we will use this class name to replace it eg. "sks woman" -> "woman"
self.diff_output_preservation_class = kwargs.get('diff_output_preservation_class', '')
# legacy
if match_adapter_assist and self.match_adapter_chance == 0.0:
self.match_adapter_chance = 1.0
# standardize inputs to the meand std of the model knowledge
self.standardize_images = kwargs.get('standardize_images', False)
self.standardize_latents = kwargs.get('standardize_latents', False)
# if self.train_turbo and not self.noise_scheduler.startswith("euler"):
# raise ValueError(f"train_turbo is only supported with euler and wuler_a noise schedulers")
self.dynamic_noise_offset = kwargs.get('dynamic_noise_offset', False)
self.do_cfg = kwargs.get('do_cfg', False)
self.do_random_cfg = kwargs.get('do_random_cfg', False)
self.cfg_scale = kwargs.get('cfg_scale', 1.0)
self.max_cfg_scale = kwargs.get('max_cfg_scale', self.cfg_scale)
self.cfg_rescale = kwargs.get('cfg_rescale', None)
if self.cfg_rescale is None:
self.cfg_rescale = self.cfg_scale
# applies the inverse of the prediction mean and std to the target to correct
# for norm drift
self.correct_pred_norm = kwargs.get('correct_pred_norm', False)
self.correct_pred_norm_multiplier = kwargs.get('correct_pred_norm_multiplier', 1.0)
self.loss_type = kwargs.get('loss_type', 'mse') # mse, mae, wavelet
# scale the prediction by this. Increase for more detail, decrease for less
self.pred_scaler = kwargs.get('pred_scaler', 1.0)
# repeats the prompt a few times to saturate the encoder
self.prompt_saturation_chance = kwargs.get('prompt_saturation_chance', 0.0)
# applies negative loss on the prior to encourage network to diverge from it
self.do_prior_divergence = kwargs.get('do_prior_divergence', False)
ema_config: Union[Dict, None] = kwargs.get('ema_config', None)
# if it is set explicitly to false, leave it false.
if ema_config is not None and ema_config.get('use_ema', None) is not None:
ema_config['use_ema'] = True
print(f"Using EMA")
else:
ema_config = {'use_ema': False}
self.ema_config: EMAConfig = EMAConfig(**ema_config)
# adds an additional loss to the network to encourage it output a normalized standard deviation
self.target_norm_std = kwargs.get('target_norm_std', None)
self.target_norm_std_value = kwargs.get('target_norm_std_value', 1.0)
self.timestep_type = kwargs.get('timestep_type', 'sigmoid') # sigmoid, linear, lognorm_blend
self.linear_timesteps = kwargs.get('linear_timesteps', False)
self.linear_timesteps2 = kwargs.get('linear_timesteps2', False)
self.disable_sampling = kwargs.get('disable_sampling', False)
# will cache a blank prompt or the trigger word, and unload the text encoder to cpu
# will make training faster and use less vram
self.unload_text_encoder = kwargs.get('unload_text_encoder', False)
# for swapping which parameters are trained during training
self.do_paramiter_swapping = kwargs.get('do_paramiter_swapping', False)
# 0.1 is 10% of the parameters active at a time lower is less vram, higher is more
self.paramiter_swapping_factor = kwargs.get('paramiter_swapping_factor', 0.1)
# bypass the guidance embedding for training. For open flux with guidance embedding
self.bypass_guidance_embedding = kwargs.get('bypass_guidance_embedding', False)
# diffusion feature extractor
self.diffusion_feature_extractor_path = kwargs.get('diffusion_feature_extractor_path', None)
self.diffusion_feature_extractor_weight = kwargs.get('diffusion_feature_extractor_weight', 1.0)
# optimal noise pairing
self.optimal_noise_pairing_samples = kwargs.get('optimal_noise_pairing_samples', 1)
# forces same noise for the same image at a given size.
self.force_consistent_noise = kwargs.get('force_consistent_noise', False)
ModelArch = Literal['sd1', 'sd2', 'sd3', 'sdxl', 'pixart', 'pixart_sigma', 'auraflow', 'flux', 'flex2', 'lumina2', 'vega', 'ssd', 'wan21']
class ModelConfig:
def __init__(self, **kwargs):
self.name_or_path: str = kwargs.get('name_or_path', None)
# name or path is updated on fine tuning. Keep a copy of the original
self.name_or_path_original: str = self.name_or_path
self.is_v2: bool = kwargs.get('is_v2', False)
self.is_xl: bool = kwargs.get('is_xl', False)
self.is_pixart: bool = kwargs.get('is_pixart', False)
self.is_pixart_sigma: bool = kwargs.get('is_pixart_sigma', False)
self.is_auraflow: bool = kwargs.get('is_auraflow', False)
self.is_v3: bool = kwargs.get('is_v3', False)
self.is_flux: bool = kwargs.get('is_flux', False)
self.is_flex2: bool = kwargs.get('is_flex2', False)
if self.is_flex2:
self.is_flux = True
self.is_lumina2: bool = kwargs.get('is_lumina2', False)
if self.is_pixart_sigma:
self.is_pixart = True
self.use_flux_cfg = kwargs.get('use_flux_cfg', False)
self.is_ssd: bool = kwargs.get('is_ssd', False)
self.is_vega: bool = kwargs.get('is_vega', False)
self.is_v_pred: bool = kwargs.get('is_v_pred', False)
self.dtype: str = kwargs.get('dtype', 'float16')
self.vae_path = kwargs.get('vae_path', None)
self.refiner_name_or_path = kwargs.get('refiner_name_or_path', None)
self._original_refiner_name_or_path = self.refiner_name_or_path
self.refiner_start_at = kwargs.get('refiner_start_at', 0.5)
self.lora_path = kwargs.get('lora_path', None)
# mainly for decompression loras for distilled models
self.assistant_lora_path = kwargs.get('assistant_lora_path', None)
self.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
self.use_text_encoder_1: bool = kwargs.get('use_text_encoder_1', True)
@@ -265,6 +504,109 @@ class ModelConfig:
# sed sdxl as true since it is mostly the same architecture
self.is_xl = True
if self.is_vega:
self.is_xl = True
# for text encoder quant. Only works with pixart currently
self.text_encoder_bits = kwargs.get('text_encoder_bits', 16) # 16, 8, 4
self.unet_path = kwargs.get("unet_path", None)
self.unet_sample_size = kwargs.get("unet_sample_size", None)
self.vae_device = kwargs.get("vae_device", None)
self.vae_dtype = kwargs.get("vae_dtype", self.dtype)
self.te_device = kwargs.get("te_device", None)
self.te_dtype = kwargs.get("te_dtype", self.dtype)
# only for flux for now
self.quantize = kwargs.get("quantize", False)
self.quantize_te = kwargs.get("quantize_te", self.quantize)
self.qtype = kwargs.get("qtype", "qfloat8")
self.qtype_te = kwargs.get("qtype_te", "qfloat8")
self.low_vram = kwargs.get("low_vram", False)
self.attn_masking = kwargs.get("attn_masking", False)
if self.attn_masking and not self.is_flux:
raise ValueError("attn_masking is only supported with flux models currently")
# for targeting a specific layers
self.ignore_if_contains: Optional[List[str]] = kwargs.get("ignore_if_contains", None)
self.only_if_contains: Optional[List[str]] = kwargs.get("only_if_contains", None)
self.quantize_kwargs = kwargs.get("quantize_kwargs", {})
# splits the model over the available gpus WIP
self.split_model_over_gpus = kwargs.get("split_model_over_gpus", False)
if self.split_model_over_gpus and not self.is_flux:
raise ValueError("split_model_over_gpus is only supported with flux models currently")
self.split_model_other_module_param_count_scale = kwargs.get("split_model_other_module_param_count_scale", 0.3)
self.te_name_or_path = kwargs.get("te_name_or_path", None)
self.arch: ModelArch = kwargs.get("arch", None)
# handle migrating to new model arch
if self.arch is not None:
# reverse the arch to the old style
if self.arch == 'sd2':
self.is_v2 = True
elif self.arch == 'sd3':
self.is_v3 = True
elif self.arch == 'sdxl':
self.is_xl = True
elif self.arch == 'pixart':
self.is_pixart = True
elif self.arch == 'pixart_sigma':
self.is_pixart_sigma = True
elif self.arch == 'auraflow':
self.is_auraflow = True
elif self.arch == 'flux':
self.is_flux = True
elif self.arch == 'flex2':
self.is_flex2 = True
elif self.arch == 'lumina2':
self.is_lumina2 = True
elif self.arch == 'vega':
self.is_vega = True
elif self.arch == 'ssd':
self.is_ssd = True
else:
pass
if self.arch is None:
if kwargs.get('is_v2', False):
self.arch = 'sd2'
elif kwargs.get('is_v3', False):
self.arch = 'sd3'
elif kwargs.get('is_xl', False):
self.arch = 'sdxl'
elif kwargs.get('is_pixart', False):
self.arch = 'pixart'
elif kwargs.get('is_pixart_sigma', False):
self.arch = 'pixart_sigma'
elif kwargs.get('is_auraflow', False):
self.arch = 'auraflow'
elif kwargs.get('is_flux', False):
self.arch = 'flux'
elif kwargs.get('is_flex2', False):
self.arch = 'flex2'
elif kwargs.get('is_lumina2', False):
self.arch = 'lumina2'
elif kwargs.get('is_vega', False):
self.arch = 'vega'
elif kwargs.get('is_ssd', False):
self.arch = 'ssd'
else:
self.arch = 'sd1'
class EMAConfig:
def __init__(self, **kwargs):
self.use_ema: bool = kwargs.get('use_ema', False)
self.ema_decay: float = kwargs.get('ema_decay', 0.999)
# feeds back the decay difference into the parameter
self.use_feedback: bool = kwargs.get('use_feedback', False)
# 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:
def __init__(self, **kwargs):
@@ -352,29 +694,46 @@ class DatasetConfig:
self.dataset_path: str = kwargs.get('dataset_path', None)
self.default_caption: str = kwargs.get('default_caption', None)
self.random_triggers: List[str] = kwargs.get('random_triggers', [])
# 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):
with open(random_triggers, 'r') as f:
random_triggers = f.read().splitlines()
# remove empty lines
random_triggers = [line for line in random_triggers if line.strip() != '']
self.random_triggers: List[str] = random_triggers
self.random_triggers_max: int = kwargs.get('random_triggers_max', 1)
self.caption_ext: str = kwargs.get('caption_ext', None)
self.random_scale: bool = kwargs.get('random_scale', False)
self.random_crop: bool = kwargs.get('random_crop', False)
self.resolution: int = kwargs.get('resolution', 512)
self.scale: float = kwargs.get('scale', 1.0)
self.buckets: bool = kwargs.get('buckets', False)
self.buckets: bool = kwargs.get('buckets', True)
self.bucket_tolerance: int = kwargs.get('bucket_tolerance', 64)
self.is_reg: bool = kwargs.get('is_reg', False)
self.network_weight: float = float(kwargs.get('network_weight', 1.0))
self.token_dropout_rate: float = float(kwargs.get('token_dropout_rate', 0.0))
self.shuffle_tokens: bool = kwargs.get('shuffle_tokens', False)
self.caption_dropout_rate: float = float(kwargs.get('caption_dropout_rate', 0.0))
self.keep_tokens: int = kwargs.get('keep_tokens', 0) # #of first tokens to always keep unless caption dropped
self.flip_x: bool = kwargs.get('flip_x', False)
self.flip_y: bool = kwargs.get('flip_y', False)
self.augments: List[str] = kwargs.get('augments', [])
self.control_path: str = kwargs.get('control_path', None) # depth maps, etc
self.control_path: Union[str,List[str]] = kwargs.get('control_path', None) # depth maps, etc
# inpaint images should be webp/png images with alpha channel. The alpha 0 (invisible) section will
# be the part conditioned to be inpainted. The alpha 1 (visible) section will be the part that is ignored
self.inpaint_path: Union[str,List[str]] = kwargs.get('inpaint_path', None)
# instead of cropping ot match image, it will serve the full size control image (clip images ie for ip adapters)
self.full_size_control_images: bool = kwargs.get('full_size_control_images', False)
self.alpha_mask: bool = kwargs.get('alpha_mask', False) # if true, will use alpha channel as mask
self.mask_path: str = kwargs.get('mask_path',
None) # focus mask (black and white. White has higher loss than black)
self.unconditional_path: str = kwargs.get('unconditional_path', None) # path where matching unconditional images are located
self.unconditional_path: str = kwargs.get('unconditional_path',
None) # path where matching unconditional images are located
self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask
self.mask_min_value: float = kwargs.get('mask_min_value', 0.01) # min value for . 0 - 1
self.mask_min_value: float = kwargs.get('mask_min_value', 0.0) # min value for . 0 - 1
self.poi: Union[str, None] = kwargs.get('poi',
None) # if one is set and in json data, will be used as auto crop scale point of interes
self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset
@@ -382,6 +741,9 @@ class DatasetConfig:
self.cache_latents: bool = kwargs.get('cache_latents', False)
# cache latents to disk will store them on disk. If both are true, it will save to disk, but keep in memory
self.cache_latents_to_disk: bool = kwargs.get('cache_latents_to_disk', False)
self.cache_clip_vision_to_disk: bool = kwargs.get('cache_clip_vision_to_disk', False)
self.standardize_images: bool = kwargs.get('standardize_images', False)
# https://albumentations.ai/docs/api_reference/augmentations/transforms
# augmentations are returned as a separate image and cannot currently be cached
@@ -400,6 +762,39 @@ class DatasetConfig:
if legacy_caption_type:
self.caption_ext = legacy_caption_type
self.caption_type = self.caption_ext
self.guidance_type: GuidanceType = kwargs.get('guidance_type', 'targeted')
# ip adapter / reference dataset
self.clip_image_path: str = kwargs.get('clip_image_path', None) # depth maps, etc
# 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', [])
self.loss_multiplier: float = kwargs.get('loss_multiplier', 1.0)
self.num_workers: int = kwargs.get('num_workers', 2)
self.prefetch_factor: int = kwargs.get('prefetch_factor', 2)
self.extra_values: List[float] = kwargs.get('extra_values', [])
self.square_crop: bool = kwargs.get('square_crop', False)
# apply same augmentations to control images. Usually want this true unless special case
self.replay_transforms: bool = kwargs.get('replay_transforms', True)
# for video
# if num_frames is greater than 1, the dataloader will look for video files.
# num_frames will be the number of frames in the training batch. If num_frames is 1, it will look for images
self.num_frames: int = kwargs.get('num_frames', 1)
# if true, will shrink video to our frames. For instance, if we have a video with 100 frames and num_frames is 10,
# we would pull frame 0, 10, 20, 30, 40, 50, 60, 70, 80, 90 so they are evenly spaced
self.shrink_video_to_frames: bool = kwargs.get('shrink_video_to_frames', True)
# fps is only used if shrink_video_to_frames is false. This will attempt to pull the num_frames at the given fps
# it will select a random start frame and pull the frames at the given fps
# this could have various issues with shorter videos and videos with variable fps
# I recommend trimming your videos to the desired length and using shrink_video_to_frames(default)
self.fps: int = kwargs.get('fps', 16)
# debug the frame count and frame selection. You dont need this. It is for debugging.
self.debug: bool = kwargs.get('debug', False)
def preprocess_dataset_raw_config(raw_config: List[dict]) -> List[dict]:
@@ -448,6 +843,11 @@ class GenerateImageConfig:
latents: Union[torch.Tensor | None] = None, # input latent to start with,
extra_kwargs: dict = None, # extra data to save with prompt file
refiner_start_at: float = 0.5, # start at this percentage of a step. 0.0 to 1.0 . 1.0 is the end
extra_values: List[float] = None, # extra values to save with prompt file
logger: Optional[EmptyLogger] = None,
num_frames: int = 1,
fps: int = 15,
ctrl_idx: int = 0
):
self.width: int = width
self.height: int = height
@@ -475,6 +875,11 @@ class GenerateImageConfig:
self.adapter_conditioning_scale: float = adapter_conditioning_scale
self.extra_kwargs = extra_kwargs if extra_kwargs is not None else {}
self.refiner_start_at = refiner_start_at
self.extra_values = extra_values if extra_values is not None else []
self.num_frames = num_frames
self.fps = fps
self.ctrl_idx = ctrl_idx
# prompt string will override any settings above
self._process_prompt_string()
@@ -484,7 +889,7 @@ class GenerateImageConfig:
self.negative_prompt_2 = negative_prompt
if prompt_2 is None:
self.prompt_2 = prompt
self.prompt_2 = self.prompt
# parse prompt paths
if self.output_path is None and self.output_folder is None:
@@ -504,6 +909,8 @@ class GenerateImageConfig:
self.height = max(64, self.height - self.height % 8) # round to divisible by 8
self.width = max(64, self.width - self.width % 8) # round to divisible by 8
self.logger = logger
def set_gen_time(self, gen_time: int = None):
if gen_time is not None:
self.gen_time = gen_time
@@ -539,11 +946,30 @@ class GenerateImageConfig:
# make parent dirs
os.makedirs(self.output_folder, exist_ok=True)
self.set_gen_time()
# TODO save image gen header info for A1111 and us, our seeds probably wont match
image.save(self.get_image_path(count, max_count))
# do prompt file
if self.add_prompt_file:
self.save_prompt_file(count, max_count)
if isinstance(image, list):
# video
if self.num_frames == 1:
raise ValueError(f"Expected 1 img but got a list {len(image)}")
if self.output_ext == 'webp':
# save as animated webp
duration = 1000 // self.fps # Convert fps to milliseconds per frame
image[0].save(
self.get_image_path(count, max_count),
format='WEBP',
append_images=image[1:],
save_all=True,
duration=duration, # Duration per frame in milliseconds
loop=0, # 0 means loop forever
quality=80 # Quality setting (0-100)
)
else:
raise ValueError(f"Unsupported video format {self.output_ext}")
else:
# TODO save image gen header info for A1111 and us, our seeds probably wont match
image.save(self.get_image_path(count, max_count))
# do prompt file
if self.add_prompt_file:
self.save_prompt_file(count, max_count)
def save_prompt_file(self, count: int = 0, max_count=0):
# save prompt file
@@ -564,7 +990,10 @@ class GenerateImageConfig:
prompt += ' --gr ' + str(self.guidance_rescale)
# get gen info
f.write(self.prompt)
try:
f.write(self.prompt)
except Exception as e:
print(f"Error writing prompt file. Prompt contains non-unicode characters. {e}")
def _process_prompt_string(self):
# we will try to support all sd-scripts where we can
@@ -633,6 +1062,18 @@ class GenerateImageConfig:
self.adapter_conditioning_scale = float(content)
elif flag == 'ref':
self.refiner_start_at = float(content)
elif flag == 'ev':
# split by comma
self.extra_values = [float(val) for val in content.split(',')]
elif flag == 'extra_values':
# split by comma
self.extra_values = [float(val) for val in content.split(',')]
elif flag == 'frames':
self.num_frames = int(content)
elif flag == 'fps':
self.fps = int(content)
elif flag == 'ctrl_idx':
self.ctrl_idx = int(content)
def post_process_embeddings(
self,
@@ -641,3 +1082,23 @@ class GenerateImageConfig:
):
# this is called after prompt embeds are encoded. We can override them in the future here
pass
def log_image(self, image, count: int = 0, max_count=0):
if self.logger is None:
return
self.logger.log_image(image, count, self.prompt)
def validate_configs(
train_config: TrainConfig,
model_config: ModelConfig,
save_config: SaveConfig,
):
if model_config.is_flux:
if save_config.save_format != 'diffusers':
# make it diffusers
save_config.save_format = 'diffusers'
if model_config.use_flux_cfg:
# bypass the embedding
train_config.bypass_guidance_embedding = True

1252
toolkit/custom_adapter.py Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -8,6 +8,7 @@ from typing import List, TYPE_CHECKING
import cv2
import numpy as np
import torch
from PIL import Image
from PIL.ImageOps import exif_transpose
from torchvision import transforms
@@ -17,11 +18,61 @@ import albumentations as A
from toolkit.buckets import get_bucket_for_image_size, BucketResolution
from toolkit.config_modules import DatasetConfig, preprocess_dataset_raw_config
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments
from toolkit.dataloader_mixins import CaptionMixin, BucketsMixin, LatentCachingMixin, Augments, CLIPCachingMixin
from toolkit.data_transfer_object.data_loader import FileItemDTO, DataLoaderBatchDTO
from toolkit.print import print_acc
from toolkit.accelerator import get_accelerator
import platform
def is_native_windows():
return platform.system() == "Windows" and platform.release() != "2"
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
image_extensions = ['.jpg', '.jpeg', '.png', '.webp']
video_extensions = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv']
class RescaleTransform:
"""Transform to rescale images to the range [-1, 1]."""
def __call__(self, image):
return image * 2 - 1
class NormalizeSDXLTransform:
"""
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
Mean: tensor([ 0.0002, -0.1034, -0.1879])
Standard Deviation: tensor([0.5436, 0.5116, 0.5033])
"""
def __call__(self, image):
return transforms.Normalize(
mean=[0.0002, -0.1034, -0.1879],
std=[0.5436, 0.5116, 0.5033],
)(image)
class NormalizeSD15Transform:
"""
Transforms the range from 0 to 1 to SDXL mean and std per channel based on avgs over thousands of images
Mean: tensor([-0.1600, -0.2450, -0.3227])
Standard Deviation: tensor([0.5319, 0.4997, 0.5139])
"""
def __call__(self, image):
return transforms.Normalize(
mean=[-0.1600, -0.2450, -0.3227],
std=[0.5319, 0.4997, 0.5139],
)(image)
class ImageDataset(Dataset, CaptionMixin):
@@ -45,7 +96,7 @@ class ImageDataset(Dataset, CaptionMixin):
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))]
# this might take a while
print(f" - Preprocessing image dimensions")
print_acc(f" - Preprocessing image dimensions")
new_file_list = []
bad_count = 0
for file in tqdm(self.file_list):
@@ -57,13 +108,13 @@ 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([
transforms.ToTensor(),
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
RescaleTransform(),
])
def get_config(self, key, default=None, required=False):
@@ -80,7 +131,13 @@ class ImageDataset(Dataset, CaptionMixin):
def __getitem__(self, index):
img_path = self.file_list[index]
img = exif_transpose(Image.open(img_path)).convert('RGB')
try:
img = exif_transpose(Image.open(img_path)).convert('RGB')
except Exception as e:
print_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))
# Downscale the source image first
img = img.resize((int(img.size[0] * self.scale), int(img.size[1] * self.scale)), Image.BICUBIC)
@@ -89,7 +146,7 @@ class ImageDataset(Dataset, CaptionMixin):
if self.random_crop:
if self.random_scale and min_img_size > self.resolution:
if min_img_size < self.resolution:
print(
print_acc(
f"Unexpected values: min_img_size={min_img_size}, self.resolution={self.resolution}, image file={img_path}")
scale_size = self.resolution
else:
@@ -192,15 +249,15 @@ 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(),
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
RescaleTransform(),
])
def get_all_prompts(self):
@@ -315,7 +372,7 @@ class PairedImageDataset(Dataset):
return img, prompt, (self.neg_weight, self.pos_weight)
class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
class AiToolkitDataset(LatentCachingMixin, CLIPCachingMixin, BucketsMixin, CaptionMixin, Dataset):
def __init__(
self,
@@ -323,8 +380,9 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
batch_size=1,
sd: 'StableDiffusion' = None,
):
super().__init__()
self.dataset_config = dataset_config
self.is_video = dataset_config.num_frames > 1
super().__init__()
folder_path = dataset_config.folder_path
self.dataset_path = dataset_config.dataset_path
if self.dataset_path is None:
@@ -333,6 +391,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
self.is_caching_latents = dataset_config.cache_latents or dataset_config.cache_latents_to_disk
self.is_caching_latents_to_memory = dataset_config.cache_latents
self.is_caching_latents_to_disk = dataset_config.cache_latents_to_disk
self.is_caching_clip_vision_to_disk = dataset_config.cache_clip_vision_to_disk
self.epoch_num = 0
self.sd = sd
@@ -353,10 +412,11 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
# check if dataset_path is a folder or json
if os.path.isdir(self.dataset_path):
file_list = [
os.path.join(self.dataset_path, file) for file in os.listdir(self.dataset_path) if
file.lower().endswith(('.jpg', '.jpeg', '.png', '.webp'))
]
extensions = image_extensions
if self.is_video:
# only look for videos
extensions = video_extensions
file_list = [os.path.join(root, file) for root, _, files in os.walk(self.dataset_path) for file in files if file.lower().endswith(tuple(extensions))]
else:
# assume json
with open(self.dataset_path, 'r') as f:
@@ -368,29 +428,88 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
# repeat the list
file_list = file_list * self.dataset_config.num_repeats
if self.dataset_config.standardize_images:
if self.sd.is_xl or self.sd.is_vega or self.sd.is_ssd:
NormalizeMethod = NormalizeSDXLTransform
else:
NormalizeMethod = NormalizeSD15Transform
self.transform = transforms.Compose([
transforms.ToTensor(),
RescaleTransform(),
NormalizeMethod(),
])
else:
self.transform = transforms.Compose([
transforms.ToTensor(),
RescaleTransform(),
])
# this might take a while
print(f" - Preprocessing image dimensions")
print_acc(f"Dataset: {self.dataset_path}")
if self.is_video:
print_acc(f" - Preprocessing video dimensions")
else:
print_acc(f" - Preprocessing image dimensions")
dataset_folder = self.dataset_path
if not os.path.isdir(self.dataset_path):
dataset_folder = os.path.dirname(dataset_folder)
dataset_size_file = os.path.join(dataset_folder, '.aitk_size.json')
dataloader_version = "0.1.1"
if os.path.exists(dataset_size_file):
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):
try:
file_item = FileItemDTO(
sd=self.sd,
path=file,
dataset_config=dataset_config
dataset_config=dataset_config,
dataloader_transforms=self.transform,
size_database=self.size_database,
dataset_root=dataset_folder,
)
self.file_list.append(file_item)
except Exception as e:
print(traceback.format_exc())
print(f"Error processing image: {file}")
print(e)
print_acc(traceback.format_exc())
if self.is_video:
print_acc(f"Error processing video: {file}")
else:
print_acc(f"Error processing image: {file}")
print_acc(e)
bad_count += 1
print(f" - Found {len(self.file_list)} images")
# print(f" - Found {bad_count} images that are too small")
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
# save the size database
with open(dataset_size_file, 'w') as f:
json.dump(self.size_database, f)
if self.is_video:
print_acc(f" - Found {len(self.file_list)} videos")
assert len(self.file_list) > 0, f"no videos found in {self.dataset_path}"
else:
print_acc(f" - Found {len(self.file_list)} images")
assert len(self.file_list) > 0, f"no images found in {self.dataset_path}"
# handle x axis flips
if self.dataset_config.flip_x:
print(" - adding x axis flips")
print_acc(" - adding x axis flips")
current_file_list = [x for x in self.file_list]
for file_item in current_file_list:
# create a copy that is flipped on the x axis
@@ -400,7 +519,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
# 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
@@ -409,12 +528,10 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
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")
self.transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize([0.5], [0.5]), # normalize to [-1, 1]
])
if self.is_video:
print_acc(f" - Found {len(self.file_list)} videos after adding flips")
else:
print_acc(f" - Found {len(self.file_list)} images after adding flips")
self.setup_epoch()
@@ -427,6 +544,8 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
self.setup_buckets()
if self.is_caching_latents:
self.cache_latents_all_latents()
if self.is_caching_clip_vision_to_disk:
self.cache_clip_vision_to_disk()
else:
if self.dataset_config.poi is not None:
# handle cropping to a specific point of interest
@@ -440,7 +559,7 @@ class AiToolkitDataset(LatentCachingMixin, BucketsMixin, CaptionMixin, Dataset):
return len(self.file_list)
def _get_single_item(self, index) -> 'FileItemDTO':
file_item = copy.deepcopy(self.file_list[index])
file_item: 'FileItemDTO' = copy.deepcopy(self.file_list[index])
file_item.load_and_process_image(self.transform)
file_item.load_caption(self.caption_dict)
return file_item
@@ -509,6 +628,13 @@ def get_dataloader_from_datasets(
# check if is caching latents
dataloader_kwargs = {}
if is_native_windows():
dataloader_kwargs['num_workers'] = 0
else:
dataloader_kwargs['num_workers'] = dataset_config_list[0].num_workers
dataloader_kwargs['prefetch_factor'] = dataset_config_list[0].prefetch_factor
if has_buckets:
# make sure they all have buckets
@@ -521,15 +647,15 @@ def get_dataloader_from_datasets(
drop_last=False,
shuffle=True,
collate_fn=dto_collation, # Use the custom collate function
num_workers=4
**dataloader_kwargs
)
else:
data_loader = DataLoader(
concatenated_dataset,
batch_size=batch_size,
shuffle=True,
num_workers=4,
collate_fn=dto_collation
collate_fn=dto_collation,
**dataloader_kwargs
)
return data_loader
@@ -556,3 +682,19 @@ def trigger_dataloader_setup_epoch(dataloader: DataLoader):
if hasattr(sub_dataset, 'setup_epoch'):
sub_dataset.setup_epoch()
sub_dataset.len = None
def get_dataloader_datasets(dataloader: DataLoader):
# hacky but needed because of different types of datasets and dataloaders
if isinstance(dataloader.dataset, list):
datasets = []
for dataset in dataloader.dataset:
if hasattr(dataset, 'datasets'):
for sub_dataset in dataset.datasets:
datasets.append(sub_dataset)
else:
datasets.append(dataset)
return datasets
elif hasattr(dataloader.dataset, 'datasets'):
return dataloader.dataset.datasets
else:
return [dataloader.dataset]

View File

@@ -1,4 +1,8 @@
import os
import weakref
from _weakref import ReferenceType
from typing import TYPE_CHECKING, List, Union
import cv2
import torch
import random
@@ -8,10 +12,12 @@ from PIL.ImageOps import exif_transpose
from toolkit import image_utils
from toolkit.dataloader_mixins import CaptionProcessingDTOMixin, ImageProcessingDTOMixin, LatentCachingFileItemDTOMixin, \
ControlFileItemDTOMixin, ArgBreakMixin, PoiFileItemDTOMixin, MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, \
UnconditionalFileItemDTOMixin
UnconditionalFileItemDTOMixin, ClipImageFileItemDTOMixin, InpaintControlFileItemDTOMixin
if TYPE_CHECKING:
from toolkit.config_modules import DatasetConfig
from toolkit.stable_diffusion_model import StableDiffusion
printed_messages = []
@@ -28,6 +34,8 @@ class FileItemDTO(
CaptionProcessingDTOMixin,
ImageProcessingDTOMixin,
ControlFileItemDTOMixin,
InpaintControlFileItemDTOMixin,
ClipImageFileItemDTOMixin,
MaskFileItemDTOMixin,
AugmentationFileItemDTOMixin,
UnconditionalFileItemDTOMixin,
@@ -35,18 +43,47 @@ class FileItemDTO(
ArgBreakMixin,
):
def __init__(self, *args, **kwargs):
self.path = kwargs.get('path', None)
self.path = kwargs.get('path', '')
self.dataset_config: 'DatasetConfig' = kwargs.get('dataset_config', None)
# process width and height
try:
w, h = image_utils.get_image_size(self.path)
except image_utils.UnknownImageFormat:
print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
f'This process is faster for png, jpeg')
self.is_video = self.dataset_config.num_frames > 1
size_database = kwargs.get('size_database', {})
dataset_root = kwargs.get('dataset_root', None)
if dataset_root is not None:
# remove dataset root from path
file_key = self.path.replace(dataset_root, '')
else:
file_key = os.path.basename(self.path)
if file_key in size_database:
w, h = size_database[file_key]
elif self.is_video:
# Open the video file
video = cv2.VideoCapture(self.path)
# Check if video opened successfully
if not video.isOpened():
raise Exception(f"Error: Could not open video file {self.path}")
# Get width and height
width = int(video.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(video.get(cv2.CAP_PROP_FRAME_HEIGHT))
# Release the video capture object immediately
video.release()
size_database[file_key] = (width, height)
else:
# original method is significantly faster, but some images are read sideways. Not sure why. Do slow method for now.
# process width and height
# try:
# w, h = image_utils.get_image_size(self.path)
# except image_utils.UnknownImageFormat:
# print_once(f'Warning: Some images in the dataset cannot be fast read. ' + \
# f'This process is faster for png, jpeg')
img = exif_transpose(Image.open(self.path))
h, w = img.size
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)
super().__init__(*args, **kwargs)
# self.caption_path: str = kwargs.get('caption_path', None)
@@ -62,6 +99,7 @@ class FileItemDTO(
self.flip_x: bool = kwargs.get('flip_x', False)
self.flip_y: bool = kwargs.get('flip_x', False)
self.augments: List[str] = self.dataset_config.augments
self.loss_multiplier: float = self.dataset_config.loss_multiplier
self.network_weight: float = self.dataset_config.network_weight
self.is_reg = self.dataset_config.is_reg
@@ -71,6 +109,8 @@ class FileItemDTO(
self.tensor = None
self.cleanup_latent()
self.cleanup_control()
self.cleanup_inpaint()
self.cleanup_clip_image()
self.cleanup_mask()
self.cleanup_unconditional()
@@ -83,11 +123,15 @@ class DataLoaderBatchDTO:
self.tensor: Union[torch.Tensor, None] = None
self.latents: Union[torch.Tensor, None] = None
self.control_tensor: Union[torch.Tensor, None] = None
self.clip_image_tensor: Union[torch.Tensor, None] = None
self.mask_tensor: Union[torch.Tensor, None] = None
self.unaugmented_tensor: Union[torch.Tensor, None] = None
self.unconditional_tensor: Union[torch.Tensor, None] = None
self.unconditional_latents: Union[torch.Tensor, None] = None
self.clip_image_embeds: Union[List[dict], None] = None
self.clip_image_embeds_unconditional: Union[List[dict], None] = None
self.sigmas: Union[torch.Tensor, None] = None # can be added elseware and passed along training code
self.extra_values: Union[torch.Tensor, None] = torch.tensor([x.extra_values for x in self.file_items]) if len(self.file_items[0].extra_values) > 0 else None
if not is_latents_cached:
# only return a tensor if latents are not cached
self.tensor: torch.Tensor = torch.cat([x.tensor.unsqueeze(0) for x in self.file_items])
@@ -112,6 +156,39 @@ class DataLoaderBatchDTO:
else:
control_tensors.append(x.control_tensor)
self.control_tensor = torch.cat([x.unsqueeze(0) for x in control_tensors])
self.inpaint_tensor: Union[torch.Tensor, None] = None
if any([x.inpaint_tensor is not None for x in self.file_items]):
# find one to use as a base
base_inpaint_tensor = None
for x in self.file_items:
if x.inpaint_tensor is not None:
base_inpaint_tensor = x.inpaint_tensor
break
inpaint_tensors = []
for x in self.file_items:
if x.inpaint_tensor is None:
inpaint_tensors.append(torch.zeros_like(base_inpaint_tensor))
else:
inpaint_tensors.append(x.inpaint_tensor)
self.inpaint_tensor = torch.cat([x.unsqueeze(0) for x in inpaint_tensors])
self.loss_multiplier_list: List[float] = [x.loss_multiplier for x in self.file_items]
if any([x.clip_image_tensor is not None for x in self.file_items]):
# find one to use as a base
base_clip_image_tensor = None
for x in self.file_items:
if x.clip_image_tensor is not None:
base_clip_image_tensor = x.clip_image_tensor
break
clip_image_tensors = []
for x in self.file_items:
if x.clip_image_tensor is None:
clip_image_tensors.append(torch.zeros_like(base_clip_image_tensor))
else:
clip_image_tensors.append(x.clip_image_tensor)
self.clip_image_tensor = torch.cat([x.unsqueeze(0) for x in clip_image_tensors])
if any([x.mask_tensor is not None for x in self.file_items]):
# find one to use as a base
@@ -159,6 +236,23 @@ class DataLoaderBatchDTO:
else:
unconditional_tensor.append(x.unconditional_tensor)
self.unconditional_tensor = torch.cat([x.unsqueeze(0) for x in unconditional_tensor])
if any([x.clip_image_embeds is not None for x in self.file_items]):
self.clip_image_embeds = []
for x in self.file_items:
if x.clip_image_embeds is not None:
self.clip_image_embeds.append(x.clip_image_embeds)
else:
raise Exception("clip_image_embeds is None for some file items")
if any([x.clip_image_embeds_unconditional is not None for x in self.file_items]):
self.clip_image_embeds_unconditional = []
for x in self.file_items:
if x.clip_image_embeds_unconditional is not None:
self.clip_image_embeds_unconditional.append(x.clip_image_embeds_unconditional)
else:
raise Exception("clip_image_embeds_unconditional is None for some file items")
except Exception as e:
print(e)
raise e
@@ -175,11 +269,7 @@ class DataLoaderBatchDTO:
to_replace_list=None,
add_if_not_present=True
):
return [x.get_caption(
trigger=trigger,
to_replace_list=to_replace_list,
add_if_not_present=add_if_not_present
) for x in self.file_items]
return [x.caption for x in self.file_items]
def get_caption_short_list(
self,
@@ -187,12 +277,7 @@ class DataLoaderBatchDTO:
to_replace_list=None,
add_if_not_present=True
):
return [x.get_caption(
trigger=trigger,
to_replace_list=to_replace_list,
add_if_not_present=add_if_not_present,
short_caption=True
) for x in self.file_items]
return [x.caption_short for x in self.file_items]
def cleanup(self):
del self.latents

File diff suppressed because it is too large Load Diff

88
toolkit/dequantize.py Normal file
View File

@@ -0,0 +1,88 @@
from functools import partial
from optimum.quanto.tensor import QTensor
import torch
def hacked_state_dict(self, *args, **kwargs):
orig_state_dict = self.orig_state_dict(*args, **kwargs)
new_state_dict = {}
for key, value in orig_state_dict.items():
if key.endswith("._scale"):
continue
if key.endswith(".input_scale"):
continue
if key.endswith(".output_scale"):
continue
if key.endswith("._data"):
key = key[:-6]
scale = orig_state_dict[key + "._scale"]
# scale is the original dtype
dtype = scale.dtype
scale = scale.float()
value = value.float()
dequantized = value * scale
# handle input and output scaling if they exist
input_scale = orig_state_dict.get(key + ".input_scale")
if input_scale is not None:
# make sure the tensor is 1.0
if input_scale.item() != 1.0:
raise ValueError("Input scale is not 1.0, cannot dequantize")
output_scale = orig_state_dict.get(key + ".output_scale")
if output_scale is not None:
# make sure the tensor is 1.0
if output_scale.item() != 1.0:
raise ValueError("Output scale is not 1.0, cannot dequantize")
new_state_dict[key] = dequantized.to('cpu', dtype=dtype)
else:
new_state_dict[key] = value
return new_state_dict
# hacks the state dict so we can dequantize before saving
def patch_dequantization_on_save(model):
model.orig_state_dict = model.state_dict
model.state_dict = partial(hacked_state_dict, model)
def dequantize_parameter(module: torch.nn.Module, param_name: str) -> bool:
"""
Convert a quantized parameter back to a regular Parameter with floating point values.
Args:
module: The module containing the parameter to unquantize
param_name: Name of the parameter to unquantize (e.g., 'weight', 'bias')
Returns:
bool: True if parameter was unquantized, False if it was already unquantized
"""
# Check if the parameter exists
if not hasattr(module, param_name):
raise AttributeError(f"Module has no parameter named '{param_name}'")
param = getattr(module, param_name)
# If it's not a parameter or not quantized, nothing to do
if not isinstance(param, torch.nn.Parameter):
raise TypeError(f"'{param_name}' is not a Parameter")
if not isinstance(param, QTensor):
return False
# Convert to float tensor while preserving device and requires_grad
with torch.no_grad():
float_tensor = param.float()
new_param = torch.nn.Parameter(
float_tensor,
requires_grad=param.requires_grad
)
# Replace the parameter
setattr(module, param_name, new_param)
return True

346
toolkit/ema.py Normal file
View File

@@ -0,0 +1,346 @@
from __future__ import division
from __future__ import unicode_literals
from typing import Iterable, Optional
import weakref
import copy
import contextlib
from toolkit.optimizers.optimizer_utils import copy_stochastic
import torch
# Partially based on:
# https://github.com/tensorflow/tensorflow/blob/r1.13/tensorflow/python/training/moving_averages.py
class ExponentialMovingAverage:
"""
Maintains (exponential) moving average of a set of parameters.
Args:
parameters: Iterable of `torch.nn.Parameter` (typically from
`model.parameters()`).
Note that EMA is computed on *all* provided parameters,
regardless of whether or not they have `requires_grad = True`;
this allows a single EMA object to be consistantly used even
if which parameters are trainable changes step to step.
If you want to some parameters in the EMA, do not pass them
to the object in the first place. For example:
ExponentialMovingAverage(
parameters=[p for p in model.parameters() if p.requires_grad],
decay=0.9
)
will ignore parameters that do not require grad.
decay: The exponential decay.
use_num_updates: Whether to use number of updates when computing
averages.
"""
def __init__(
self,
parameters: Iterable[torch.nn.Parameter] = None,
decay: float = 0.995,
use_num_updates: bool = False,
# feeds back the decat to the parameter
use_feedback: bool = False,
param_multiplier: float = 1.0
):
if parameters is None:
raise ValueError("parameters must be provided")
if decay < 0.0 or decay > 1.0:
raise ValueError('Decay must be between 0 and 1')
self.decay = decay
self.num_updates = 0 if use_num_updates else None
self.use_feedback = use_feedback
self.param_multiplier = param_multiplier
parameters = list(parameters)
self.shadow_params = [
p.clone().detach()
for p in parameters
]
self.collected_params = None
self._is_train_mode = True
# By maintaining only a weakref to each parameter,
# we maintain the old GC behaviour of ExponentialMovingAverage:
# if the model goes out of scope but the ExponentialMovingAverage
# is kept, no references to the model or its parameters will be
# maintained, and the model will be cleaned up.
self._params_refs = [weakref.ref(p) for p in parameters]
def _get_parameters(
self,
parameters: Optional[Iterable[torch.nn.Parameter]]
) -> Iterable[torch.nn.Parameter]:
if parameters is None:
parameters = [p() for p in self._params_refs]
if any(p is None for p in parameters):
raise ValueError(
"(One of) the parameters with which this "
"ExponentialMovingAverage "
"was initialized no longer exists (was garbage collected);"
" please either provide `parameters` explicitly or keep "
"the model to which they belong from being garbage "
"collected."
)
return parameters
else:
parameters = list(parameters)
if len(parameters) != len(self.shadow_params):
raise ValueError(
"Number of parameters passed as argument is different "
"from number of shadow parameters maintained by this "
"ExponentialMovingAverage"
)
return parameters
def update(
self,
parameters: Optional[Iterable[torch.nn.Parameter]] = None
) -> None:
"""
Update currently maintained parameters.
Call this every time the parameters are updated, such as the result of
the `optimizer.step()` call.
Args:
parameters: Iterable of `torch.nn.Parameter`; usually the same set of
parameters used to initialize this object. If `None`, the
parameters with which this `ExponentialMovingAverage` was
initialized will be used.
"""
parameters = self._get_parameters(parameters)
decay = self.decay
if self.num_updates is not None:
self.num_updates += 1
decay = min(
decay,
(1 + self.num_updates) / (10 + self.num_updates)
)
one_minus_decay = 1.0 - decay
with torch.no_grad():
for s_param, param in zip(self.shadow_params, parameters):
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_float.sub_(tmp)
update_param = False
if self.use_feedback:
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,
parameters: Optional[Iterable[torch.nn.Parameter]] = None
) -> None:
"""
Copy current averaged parameters into given collection of parameters.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
updated with the stored moving averages. If `None`, the
parameters with which this `ExponentialMovingAverage` was
initialized will be used.
"""
parameters = self._get_parameters(parameters)
for s_param, param in zip(self.shadow_params, parameters):
param.data.copy_(s_param.data)
def store(
self,
parameters: Optional[Iterable[torch.nn.Parameter]] = None
) -> None:
"""
Save the current parameters for restoring later.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
temporarily stored. If `None`, the parameters of with which this
`ExponentialMovingAverage` was initialized will be used.
"""
parameters = self._get_parameters(parameters)
self.collected_params = [
param.clone()
for param in parameters
]
def restore(
self,
parameters: Optional[Iterable[torch.nn.Parameter]] = None
) -> None:
"""
Restore the parameters stored with the `store` method.
Useful to validate the model with EMA parameters without affecting the
original optimization process. Store the parameters before the
`copy_to` method. After validation (or model saving), use this to
restore the former parameters.
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
updated with the stored parameters. If `None`, the
parameters with which this `ExponentialMovingAverage` was
initialized will be used.
"""
if self.collected_params is None:
raise RuntimeError(
"This ExponentialMovingAverage has no `store()`ed weights "
"to `restore()`"
)
parameters = self._get_parameters(parameters)
for c_param, param in zip(self.collected_params, parameters):
param.data.copy_(c_param.data)
@contextlib.contextmanager
def average_parameters(
self,
parameters: Optional[Iterable[torch.nn.Parameter]] = None
):
r"""
Context manager for validation/inference with averaged parameters.
Equivalent to:
ema.store()
ema.copy_to()
try:
...
finally:
ema.restore()
Args:
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
updated with the stored parameters. If `None`, the
parameters with which this `ExponentialMovingAverage` was
initialized will be used.
"""
parameters = self._get_parameters(parameters)
self.store(parameters)
self.copy_to(parameters)
try:
yield
finally:
self.restore(parameters)
def to(self, device=None, dtype=None) -> None:
r"""Move internal buffers of the ExponentialMovingAverage to `device`.
Args:
device: like `device` argument to `torch.Tensor.to`
"""
# .to() on the tensors handles None correctly
self.shadow_params = [
p.to(device=device, dtype=dtype)
if p.is_floating_point()
else p.to(device=device)
for p in self.shadow_params
]
if self.collected_params is not None:
self.collected_params = [
p.to(device=device, dtype=dtype)
if p.is_floating_point()
else p.to(device=device)
for p in self.collected_params
]
return
def state_dict(self) -> dict:
r"""Returns the state of the ExponentialMovingAverage as a dict."""
# Following PyTorch conventions, references to tensors are returned:
# "returns a reference to the state and not its copy!" -
# https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict
return {
"decay": self.decay,
"num_updates": self.num_updates,
"shadow_params": self.shadow_params,
"collected_params": self.collected_params
}
def load_state_dict(self, state_dict: dict) -> None:
r"""Loads the ExponentialMovingAverage state.
Args:
state_dict (dict): EMA state. Should be an object returned
from a call to :meth:`state_dict`.
"""
# deepcopy, to be consistent with module API
state_dict = copy.deepcopy(state_dict)
self.decay = state_dict["decay"]
if self.decay < 0.0 or self.decay > 1.0:
raise ValueError('Decay must be between 0 and 1')
self.num_updates = state_dict["num_updates"]
assert self.num_updates is None or isinstance(self.num_updates, int), \
"Invalid num_updates"
self.shadow_params = state_dict["shadow_params"]
assert isinstance(self.shadow_params, list), \
"shadow_params must be a list"
assert all(
isinstance(p, torch.Tensor) for p in self.shadow_params
), "shadow_params must all be Tensors"
self.collected_params = state_dict["collected_params"]
if self.collected_params is not None:
assert isinstance(self.collected_params, list), \
"collected_params must be a list"
assert all(
isinstance(p, torch.Tensor) for p in self.collected_params
), "collected_params must all be Tensors"
assert len(self.collected_params) == len(self.shadow_params), \
"collected_params and shadow_params had different lengths"
if len(self.shadow_params) == len(self._params_refs):
# Consistant with torch.optim.Optimizer, cast things to consistant
# device and dtype with the parameters
params = [p() for p in self._params_refs]
# If parameters have been garbage collected, just load the state
# we were given without change.
if not any(p is None for p in params):
# ^ parameter references are still good
for i, p in enumerate(params):
self.shadow_params[i] = self.shadow_params[i].to(
device=p.device, dtype=p.dtype
)
if self.collected_params is not None:
self.collected_params[i] = self.collected_params[i].to(
device=p.device, dtype=p.dtype
)
else:
raise ValueError(
"Tried to `load_state_dict()` with the wrong number of "
"parameters in the saved state."
)
def eval(self):
if self._is_train_mode:
with torch.no_grad():
self.store()
self.copy_to()
self._is_train_mode = False
def train(self):
if not self._is_train_mode:
with torch.no_grad():
self.restore()
self._is_train_mode = True

View File

@@ -86,18 +86,19 @@ class Embedding:
self.orig_embeds_params = [x.get_input_embeddings().weight.data.clone() for x in self.text_encoder_list]
def restore_embeddings(self):
# Let's make sure we don't update any embedding weights besides the newly added token
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
self.tokenizer_list,
self.orig_embeds_params,
self.placeholder_token_ids):
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
index_no_updates[
min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
with torch.no_grad():
with torch.no_grad():
# Let's make sure we don't update any embedding weights besides the newly added token
for text_encoder, tokenizer, orig_embeds, placeholder_token_ids in zip(self.text_encoder_list,
self.tokenizer_list,
self.orig_embeds_params,
self.placeholder_token_ids):
index_no_updates = torch.ones((len(tokenizer),), dtype=torch.bool)
index_no_updates[ min(placeholder_token_ids): max(placeholder_token_ids) + 1] = False
text_encoder.get_input_embeddings().weight[
index_no_updates
] = orig_embeds[index_no_updates]
weight = text_encoder.get_input_embeddings().weight
pass
def get_trainable_params(self):
params = []

827
toolkit/guidance.py Normal file
View File

@@ -0,0 +1,827 @@
import torch
from typing import Literal, Optional
from toolkit.basic import value_map
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
from toolkit.stable_diffusion_model import StableDiffusion
from toolkit.train_tools import get_torch_dtype
from toolkit.config_modules import TrainConfig
GuidanceType = Literal["targeted", "polarity", "targeted_polarity", "direct"]
DIFFERENTIAL_SCALER = 0.2
# DIFFERENTIAL_SCALER = 0.25
def get_differential_mask(
conditional_latents: torch.Tensor,
unconditional_latents: torch.Tensor,
threshold: float = 0.2,
gradient: bool = False,
):
# make a differential mask
differential_mask = torch.abs(conditional_latents - unconditional_latents)
if len(differential_mask.shape) == 4:
max_differential = \
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
elif len(differential_mask.shape) == 5:
max_differential = \
differential_mask.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0].max(dim=4, keepdim=True)[0]
differential_scaler = 1.0 / max_differential
differential_mask = differential_mask * differential_scaler
if gradient:
# wew need to scale it to 0-1
# differential_mask = differential_mask - differential_mask.min()
# differential_mask = differential_mask / differential_mask.max()
# add 0.2 threshold to both sides and clip
differential_mask = value_map(
differential_mask,
differential_mask.min(),
differential_mask.max(),
0 - threshold,
1 + threshold
)
differential_mask = torch.clamp(differential_mask, 0.0, 1.0)
else:
# make everything less than 0.2 be 0.0 and everything else be 1.0
differential_mask = torch.where(
differential_mask < threshold,
torch.zeros_like(differential_mask),
torch.ones_like(differential_mask)
)
return differential_mask
def get_targeted_polarity_loss(
noisy_latents: torch.Tensor,
conditional_embeds: PromptEmbeds,
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
**kwargs
):
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
with torch.no_grad():
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
# inputs_abs_mean = torch.abs(conditional_latents).mean(dim=[1, 2, 3], keepdim=True)
# noise_abs_mean = torch.abs(noise).mean(dim=[1, 2, 3], keepdim=True)
differential_scaler = DIFFERENTIAL_SCALER
unconditional_diff = (unconditional_latents - conditional_latents)
unconditional_diff_noise = unconditional_diff * differential_scaler
conditional_diff = (conditional_latents - unconditional_latents)
conditional_diff_noise = conditional_diff * differential_scaler
conditional_diff_noise = conditional_diff_noise.detach().requires_grad_(False)
unconditional_diff_noise = unconditional_diff_noise.detach().requires_grad_(False)
#
baseline_conditional_noisy_latents = sd.add_noise(
conditional_latents,
noise,
timesteps
).detach()
baseline_unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
noise,
timesteps
).detach()
conditional_noise = noise + unconditional_diff_noise
unconditional_noise = noise + conditional_diff_noise
conditional_noisy_latents = sd.add_noise(
conditional_latents,
conditional_noise,
timesteps
).detach()
unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
unconditional_noise,
timesteps
).detach()
# double up everything to run it through all at once
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
# cat_baseline_noisy_latents = torch.cat(
# [baseline_conditional_noisy_latents, baseline_unconditional_noisy_latents],
# dim=0
# )
# Disable the LoRA network so we can predict parent network knowledge without it
# sd.network.is_active = False
# sd.unet.eval()
# Predict noise to get a baseline of what the parent network wants to do with the latents + noise.
# This acts as our control to preserve the unaltered parts of the image.
# baseline_prediction = sd.predict_noise(
# latents=cat_baseline_noisy_latents.to(device, dtype=dtype).detach(),
# conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
# timestep=cat_timesteps,
# guidance_scale=1.0,
# **pred_kwargs # adapter residuals in here
# ).detach()
# conditional_baseline_prediction, unconditional_baseline_prediction = torch.chunk(baseline_prediction, 2, dim=0)
# negative_network_weights = [weight * -1.0 for weight in network_weight_list]
# positive_network_weights = [weight * 1.0 for weight in network_weight_list]
# cat_network_weight_list = positive_network_weights + negative_network_weights
# turn the LoRA network back on.
sd.unet.train()
# sd.network.is_active = True
# sd.network.multiplier = cat_network_weight_list
# do our prediction with LoRA active on the scaled guidance latents
prediction = sd.predict_noise(
latents=cat_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
timestep=cat_timesteps,
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
)
# prediction = prediction - baseline_prediction
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
# pred_pos = pred_pos - conditional_baseline_prediction
# pred_neg = pred_neg - unconditional_baseline_prediction
pred_loss = torch.nn.functional.mse_loss(
pred_pos.float(),
conditional_noise.float(),
reduction="none"
)
pred_loss = pred_loss.mean([1, 2, 3])
pred_neg_loss = torch.nn.functional.mse_loss(
pred_neg.float(),
unconditional_noise.float(),
reduction="none"
)
pred_neg_loss = pred_neg_loss.mean([1, 2, 3])
loss = pred_loss + pred_neg_loss
loss = loss.mean()
loss.backward()
# detach it so parent class can run backward on no grads without throwing error
loss = loss.detach()
loss.requires_grad_(True)
return loss
def get_direct_guidance_loss(
noisy_latents: torch.Tensor,
conditional_embeds: 'PromptEmbeds',
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
unconditional_embeds: Optional[PromptEmbeds] = None,
mask_multiplier=None,
prior_pred=None,
**kwargs
):
with torch.no_grad():
# Perform targeted guidance (working title)
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
conditional_noisy_latents = sd.add_noise(
conditional_latents,
# target_noise,
noise,
timesteps
).detach()
unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
noise,
timesteps
).detach()
# turn the LoRA network back on.
sd.unet.train()
# sd.network.is_active = True
# sd.network.multiplier = network_weight_list
# do our prediction with LoRA active on the scaled guidance latents
if unconditional_embeds is not None:
unconditional_embeds = unconditional_embeds.to(device, dtype=dtype).detach()
unconditional_embeds = concat_prompt_embeds([unconditional_embeds, unconditional_embeds])
prediction = sd.predict_noise(
latents=torch.cat([unconditional_noisy_latents, conditional_noisy_latents]).to(device, dtype=dtype).detach(),
conditional_embeddings=concat_prompt_embeds([conditional_embeds,conditional_embeds]).to(device, dtype=dtype).detach(),
unconditional_embeddings=unconditional_embeds,
timestep=torch.cat([timesteps, timesteps]),
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
)
noise_pred_uncond, noise_pred_cond = torch.chunk(prediction, 2, dim=0)
guidance_scale = 1.1
guidance_pred = noise_pred_uncond + guidance_scale * (
noise_pred_cond - noise_pred_uncond
)
guidance_loss = torch.nn.functional.mse_loss(
guidance_pred.float(),
noise.detach().float(),
reduction="none"
)
if mask_multiplier is not None:
guidance_loss = guidance_loss * mask_multiplier
guidance_loss = guidance_loss.mean([1, 2, 3])
guidance_loss = guidance_loss.mean()
# loss = guidance_loss + masked_noise_loss
loss = guidance_loss
loss.backward()
# detach it so parent class can run backward on no grads without throwing error
loss = loss.detach()
loss.requires_grad_(True)
return loss
# targeted
def get_targeted_guidance_loss(
noisy_latents: torch.Tensor,
conditional_embeds: 'PromptEmbeds',
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
**kwargs
):
with torch.no_grad():
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
# Encode the unconditional image into latents
unconditional_noisy_latents = sd.noise_scheduler.add_noise(
unconditional_latents,
noise,
timesteps
)
conditional_noisy_latents = sd.noise_scheduler.add_noise(
conditional_latents,
noise,
timesteps
)
# was_network_active = self.network.is_active
sd.network.is_active = False
sd.unet.eval()
target_differential = unconditional_latents - conditional_latents
# scale our loss by the differential scaler
target_differential_abs = target_differential.abs()
target_differential_abs_min = \
target_differential_abs.min(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
target_differential_abs_max = \
target_differential_abs.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0]
min_guidance = 1.0
max_guidance = 2.0
differential_scaler = value_map(
target_differential_abs,
target_differential_abs_min,
target_differential_abs_max,
min_guidance,
max_guidance
).detach()
# With LoRA network bypassed, predict noise to get a baseline of what the network
# wants to do with the latents + noise. Pass our target latents here for the input.
target_unconditional = sd.predict_noise(
latents=unconditional_noisy_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
timestep=timesteps,
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
).detach()
prior_prediction_loss = torch.nn.functional.mse_loss(
target_unconditional.float(),
noise.float(),
reduction="none"
).detach().clone()
# turn the LoRA network back on.
sd.unet.train()
sd.network.is_active = True
sd.network.multiplier = network_weight_list + [x + -1.0 for x in network_weight_list]
# with LoRA active, predict the noise with the scaled differential latents added. This will allow us
# the opportunity to predict the differential + noise that was added to the latents.
prediction = sd.predict_noise(
latents=torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0).to(device, dtype=dtype).detach(),
conditional_embeddings=concat_prompt_embeds([conditional_embeds, conditional_embeds]).to(device, dtype=dtype).detach(),
timestep=torch.cat([timesteps, timesteps], dim=0),
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
)
prediction_conditional, prediction_unconditional = torch.chunk(prediction, 2, dim=0)
conditional_loss = torch.nn.functional.mse_loss(
prediction_conditional.float(),
noise.float(),
reduction="none"
)
unconditional_loss = torch.nn.functional.mse_loss(
prediction_unconditional.float(),
noise.float(),
reduction="none"
)
positive_loss = torch.abs(
conditional_loss.float() - prior_prediction_loss.float(),
)
# scale our loss by the differential scaler
positive_loss = positive_loss * differential_scaler
positive_loss = positive_loss.mean([1, 2, 3])
polar_loss = torch.abs(
conditional_loss.float() - unconditional_loss.float(),
).mean([1, 2, 3])
positive_loss = positive_loss.mean() + polar_loss.mean()
positive_loss.backward()
# loss = positive_loss.detach() + negative_loss.detach()
loss = positive_loss.detach()
# add a grad so other backward does not fail
loss.requires_grad_(True)
# restore network
sd.network.multiplier = network_weight_list
return loss
def get_guided_loss_polarity(
noisy_latents: torch.Tensor,
conditional_embeds: PromptEmbeds,
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
train_config: 'TrainConfig',
scaler=None,
**kwargs
):
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
with torch.no_grad():
dtype = get_torch_dtype(dtype)
noise = noise.to(device, dtype=dtype).detach()
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
target_pos = noise
target_neg = noise
if sd.is_flow_matching:
linear_timesteps = any([
train_config.linear_timesteps,
train_config.linear_timesteps2,
train_config.timestep_type == 'linear',
])
timestep_type = 'linear' if linear_timesteps else None
if timestep_type is None:
timestep_type = train_config.timestep_type
sd.noise_scheduler.set_train_timesteps(
1000,
device=device,
timestep_type=timestep_type,
latents=conditional_latents
)
target_pos = (noise - conditional_latents).detach()
target_neg = (noise - unconditional_latents).detach()
conditional_noisy_latents = sd.add_noise(
conditional_latents,
noise,
timesteps
).detach()
unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
noise,
timesteps
).detach()
# double up everything to run it through all at once
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
negative_network_weights = [weight * -1.0 for weight in network_weight_list]
positive_network_weights = [weight * 1.0 for weight in network_weight_list]
cat_network_weight_list = positive_network_weights + negative_network_weights
# turn the LoRA network back on.
sd.unet.train()
sd.network.is_active = True
sd.network.multiplier = cat_network_weight_list
# do our prediction with LoRA active on the scaled guidance latents
prediction = sd.predict_noise(
latents=cat_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
timestep=cat_timesteps,
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
)
pred_pos, pred_neg = torch.chunk(prediction, 2, dim=0)
pred_loss = torch.nn.functional.mse_loss(
pred_pos.float(),
target_pos.float(),
reduction="none"
)
# pred_loss = pred_loss.mean([1, 2, 3])
pred_neg_loss = torch.nn.functional.mse_loss(
pred_neg.float(),
target_neg.float(),
reduction="none"
)
loss = pred_loss + pred_neg_loss
loss = loss.mean([1, 2, 3])
loss = loss.mean()
if scaler is not None:
scaler.scale(loss).backward()
else:
loss.backward()
# detach it so parent class can run backward on no grads without throwing error
loss = loss.detach()
loss.requires_grad_(True)
return loss
def get_guided_tnt(
noisy_latents: torch.Tensor,
conditional_embeds: PromptEmbeds,
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
prior_pred: torch.Tensor = None,
**kwargs
):
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
with torch.no_grad():
dtype = get_torch_dtype(dtype)
noise = noise.to(device, dtype=dtype).detach()
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
conditional_noisy_latents = sd.add_noise(
conditional_latents,
noise,
timesteps
).detach()
unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
noise,
timesteps
).detach()
# double up everything to run it through all at once
cat_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
cat_latents = torch.cat([conditional_noisy_latents, unconditional_noisy_latents], dim=0)
cat_timesteps = torch.cat([timesteps, timesteps], dim=0)
# turn the LoRA network back on.
sd.unet.train()
if sd.network is not None:
cat_network_weight_list = [weight for weight in network_weight_list * 2]
sd.network.multiplier = cat_network_weight_list
sd.network.is_active = True
prediction = sd.predict_noise(
latents=cat_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=cat_embeds.to(device, dtype=dtype).detach(),
timestep=cat_timesteps,
guidance_scale=1.0,
**pred_kwargs # adapter residuals in here
)
this_prediction, that_prediction = torch.chunk(prediction, 2, dim=0)
this_loss = torch.nn.functional.mse_loss(
this_prediction.float(),
noise.float(),
reduction="none"
)
that_loss = torch.nn.functional.mse_loss(
that_prediction.float(),
noise.float(),
reduction="none"
)
this_loss = this_loss.mean([1, 2, 3])
# negative loss on that
that_loss = -that_loss.mean([1, 2, 3])
with torch.no_grad():
# match that loss with this loss so it is not a negative value and same scale
that_loss_scaler = torch.abs(this_loss) / torch.abs(that_loss)
that_loss = that_loss * that_loss_scaler * 0.01
loss = this_loss + that_loss
loss = loss.mean()
loss.backward()
# detach it so parent class can run backward on no grads without throwing error
loss = loss.detach()
loss.requires_grad_(True)
return loss
def targeted_flow_guidance(
noisy_latents: torch.Tensor,
conditional_embeds: 'PromptEmbeds',
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
unconditional_embeds: Optional[PromptEmbeds] = None,
mask_multiplier=None,
prior_pred=None,
scaler=None,
train_config=None,
**kwargs
):
if not sd.is_flow_matching:
raise ValueError("targeted_flow only works on flow matching models")
dtype = get_torch_dtype(sd.torch_dtype)
device = sd.device_torch
with torch.no_grad():
dtype = get_torch_dtype(dtype)
noise = noise.to(device, dtype=dtype).detach()
conditional_latents = batch.latents.to(device, dtype=dtype).detach()
unconditional_latents = batch.unconditional_latents.to(device, dtype=dtype).detach()
# get a mask on the differential of the latents
# this will be scaled from 0.0-1.0 with 1.0 being the largest differential
abs_differential_mask = get_differential_mask(
conditional_latents,
unconditional_latents,
gradient=True
)
# get noisy latents for both conditional and unconditional predictions
unconditional_noisy_latents = sd.add_noise(
unconditional_latents,
noise,
timesteps
).detach()
conditional_noisy_latents = sd.add_noise(
conditional_latents,
noise,
timesteps
).detach()
# disable the lora to get a baseline prediction
sd.network.is_active = False
sd.unet.eval()
# get a baseline prediction of the model knowledge without the lora network
# we do this with the unconditional noisy latents
baseline_prediction = sd.predict_noise(
latents=unconditional_noisy_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
timestep=timesteps,
guidance_scale=1.0,
**pred_kwargs
).detach()
# This is our normal flowmatching target
# target = noise - latents
# we need to target the baseline noise but with our conditional latents
# to do this we first have to determine the baseline_prediction noise by reversing the flowmatching target
baseline_predicted_noise = baseline_prediction + unconditional_latents
# baseline_predicted_noise is now the noise prediction our model would make with a the unconditional image.
# we use this as our new noise target to preserve the existing knowledge of the image.
# we apply a mask to this noise to only allow the differential of the conditional latents to be learned
baseline_predicted_noise = (1 - abs_differential_mask) * baseline_predicted_noise
masked_noise = abs_differential_mask * noise
target_noise = masked_noise + baseline_predicted_noise
# compute our new target prediction using our current knowledge noise with our conditional latents
# this makes it so the only new information is the differential of our conditional and unconditional latents
# forcing the network to preserve existing knowledge, but learn only our changes
target_pred = (target_noise - conditional_latents).detach()
# make a prediction with the lora network active
sd.unet.train()
sd.network.is_active = True
sd.network.multiplier = network_weight_list
prediction = sd.predict_noise(
latents=conditional_noisy_latents.to(device, dtype=dtype).detach(),
conditional_embeddings=conditional_embeds.to(device, dtype=dtype).detach(),
timestep=timesteps,
guidance_scale=1.0,
**pred_kwargs
)
# target our baseline + diffirential noise target
pred_loss = torch.nn.functional.mse_loss(
prediction.float(),
target_pred.float()
)
return pred_loss
# this processes all guidance losses based on the batch information
def get_guidance_loss(
noisy_latents: torch.Tensor,
conditional_embeds: 'PromptEmbeds',
match_adapter_assist: bool,
network_weight_list: list,
timesteps: torch.Tensor,
pred_kwargs: dict,
batch: 'DataLoaderBatchDTO',
noise: torch.Tensor,
sd: 'StableDiffusion',
unconditional_embeds: Optional[PromptEmbeds] = None,
mask_multiplier=None,
prior_pred=None,
scaler=None,
train_config=None,
**kwargs
):
# TODO add others and process individual batch items separately
guidance_type: GuidanceType = batch.file_items[0].dataset_config.guidance_type
if guidance_type == "targeted":
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted guidance"
return get_targeted_guidance_loss(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
**kwargs
)
elif guidance_type == "polarity":
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
return get_guided_loss_polarity(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
scaler=scaler,
train_config=train_config,
**kwargs
)
elif guidance_type == "tnt":
assert unconditional_embeds is None, "Unconditional embeds are not supported for polarity guidance"
return get_guided_tnt(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
prior_pred=prior_pred,
**kwargs
)
elif guidance_type == "targeted_polarity":
assert unconditional_embeds is None, "Unconditional embeds are not supported for targeted polarity guidance"
return get_targeted_polarity_loss(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
**kwargs
)
elif guidance_type == "direct":
return get_direct_guidance_loss(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
unconditional_embeds=unconditional_embeds,
mask_multiplier=mask_multiplier,
prior_pred=prior_pred,
**kwargs
)
elif guidance_type == "targeted_flow":
return targeted_flow_guidance(
noisy_latents,
conditional_embeds,
match_adapter_assist,
network_weight_list,
timesteps,
pred_kwargs,
batch,
noise,
sd,
unconditional_embeds=unconditional_embeds,
mask_multiplier=mask_multiplier,
prior_pred=prior_pred,
scaler=scaler,
train_config=train_config,
**kwargs
)
else:
raise NotImplementedError(f"Guidance type {guidance_type} is not implemented")

View File

@@ -5,12 +5,14 @@ import json
import os
import io
import struct
import threading
from typing import TYPE_CHECKING
import cv2
import numpy as np
import torch
from diffusers import AutoencoderTiny
from PIL import Image as PILImage
FILE_UNKNOWN = "Sorry, don't know how to get size for this file."
@@ -425,43 +427,82 @@ def main(argv=None):
is_window_shown = False
display_lock = threading.Lock()
current_img = None
update_event = threading.Event()
def update_image(img, name):
global current_img
with display_lock:
current_img = (img, name)
update_event.set()
def display_image_in_thread():
global is_window_shown
def display_img():
global current_img
while True:
update_event.wait()
with display_lock:
if current_img:
img, name = current_img
cv2.imshow(name, img)
current_img = None
update_event.clear()
if cv2.waitKey(1) & 0xFF == 27: # Esc key to stop
cv2.destroyAllWindows()
print('\nESC pressed, stopping')
break
if not is_window_shown:
is_window_shown = True
threading.Thread(target=display_img, daemon=True).start()
def show_img(img, name='AI Toolkit'):
global is_window_shown
img = np.clip(img, 0, 255).astype(np.uint8)
cv2.imshow(name, img[:, :, ::-1])
k = cv2.waitKey(10) & 0xFF
if k == 27: # Esc key to stop
print('\nESC pressed, stopping')
raise KeyboardInterrupt
update_image(img[:, :, ::-1], name)
if not is_window_shown:
is_window_shown = True
display_image_in_thread()
def show_tensors(imgs: torch.Tensor, name='AI Toolkit'):
# if rank is 4
if len(imgs.shape) == 4:
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
else:
img_list = [imgs]
# put images side by side
img = torch.cat(img_list, dim=3)
# img is -1 to 1, convert to 0 to 255
img = img / 2 + 0.5
img_numpy = img.to(torch.float32).detach().cpu().numpy()
img_numpy = np.clip(img_numpy, 0, 1) * 255
# convert to numpy Move channel to last
img_numpy = img_numpy.transpose(0, 2, 3, 1)
# convert to uint8
img_numpy = img_numpy.astype(np.uint8)
show_img(img_numpy[0], name=name)
show_img(img_numpy[0], name=name)
def save_tensors(imgs: torch.Tensor, path='output.png'):
if len(imgs.shape) == 5 and imgs.shape[0] == 1:
imgs = imgs.squeeze(0)
if len(imgs.shape) == 4:
img_list = torch.chunk(imgs, imgs.shape[0], dim=0)
else:
img_list = [imgs]
img = torch.cat(img_list, dim=3)
img = img / 2 + 0.5
img_numpy = img.to(torch.float32).detach().cpu().numpy()
img_numpy = np.clip(img_numpy, 0, 1) * 255
img_numpy = img_numpy.transpose(0, 2, 3, 1)
img_numpy = img_numpy.astype(np.uint8)
# concat images to one
img_numpy = np.concatenate(img_numpy, axis=1)
# conver to pil
img_pil = PILImage.fromarray(img_numpy)
img_pil.save(path)
def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit'):
# decode latents
if vae.device == 'cpu':
vae.to(latents.device)
latents = latents / vae.config['scaling_factor']
@@ -469,12 +510,24 @@ def show_latents(latents: torch.Tensor, vae: 'AutoencoderTiny', name='AI Toolkit
show_tensors(imgs, name=name)
def on_exit():
if is_window_shown:
cv2.destroyAllWindows()
def reduce_contrast(tensor, factor):
# Ensure factor is between 0 and 1
factor = max(0, min(factor, 1))
# Calculate the mean of the tensor
mean = torch.mean(tensor)
# Reduce contrast
adjusted_tensor = (tensor - mean) * factor + mean
# Clip values to ensure they stay within -1 to 1 range
return torch.clamp(adjusted_tensor, -1.0, 1.0)
atexit.register(on_exit)
if __name__ == "__main__":

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

84
toolkit/logging.py Normal file
View File

@@ -0,0 +1,84 @@
from typing import OrderedDict, Optional
from PIL import Image
from toolkit.config_modules import LoggingConfig
# Base logger class
# This class does nothing, it's just a placeholder
class EmptyLogger:
def __init__(self, *args, **kwargs) -> None:
pass
# start logging the training
def start(self):
pass
# collect the log to send
def log(self, *args, **kwargs):
pass
# send the log
def commit(self, step: Optional[int] = None):
pass
# log image
def log_image(self, *args, **kwargs):
pass
# finish logging
def finish(self):
pass
# Wandb logger class
# This class logs the data to wandb
class WandbLogger(EmptyLogger):
def __init__(self, project: str, run_name: str | None, config: OrderedDict) -> None:
self.project = project
self.run_name = run_name
self.config = config
def start(self):
try:
import wandb
except ImportError:
raise ImportError("Failed to import wandb. Please install wandb by running `pip install wandb`")
# send the whole config to wandb
run = wandb.init(project=self.project, name=self.run_name, config=self.config)
self.run = run
self._log = wandb.log # log function
self._image = wandb.Image # image object
def log(self, *args, **kwargs):
# when commit is False, wandb increments the step,
# but we don't want that to happen, so we set commit=False
self._log(*args, **kwargs, commit=False)
def commit(self, step: Optional[int] = None):
# after overall one step is done, we commit the log
# by log empty object with commit=True
self._log({}, step=step, commit=True)
def log_image(
self,
image: Image,
id, # sample index
caption: str | None = None, # positive prompt
*args,
**kwargs,
):
# create a wandb image object and log it
image = self._image(image, caption=caption, *args, **kwargs)
self._log({f"sample_{id}": image}, commit=False)
def finish(self):
self.run.finish()
# create logger based on the logging config
def create_logger(logging_config: LoggingConfig, all_config: OrderedDict):
if logging_config.use_wandb:
project_name = logging_config.project_name
run_name = logging_config.run_name
return WandbLogger(project=project_name, run_name=run_name, config=all_config)
else:
return EmptyLogger()

View File

@@ -1,11 +1,15 @@
import copy
import json
import math
import weakref
import os
import re
import sys
from typing import List, Optional, Dict, Type, Union
import torch
from diffusers import UNet2DConditionModel, PixArtTransformer2DModel, AuraFlowTransformer2DModel
from transformers import CLIPTextModel
from toolkit.models.lokr import LokrModule
from .config_modules import NetworkConfig
from .lorm import count_parameters
@@ -15,21 +19,28 @@ from .paths import SD_SCRIPTS_ROOT
sys.path.append(SD_SCRIPTS_ROOT)
from networks.lora import LoRANetwork, get_block_index
from toolkit.models.DoRA import DoRAModule
from typing import TYPE_CHECKING
from torch.utils.checkpoint import checkpoint
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_")
# diffusers specific stuff
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear'
'LoRACompatibleLinear',
'QLinear',
# 'GroupNorm',
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv'
'LoRACompatibleConv',
'QConv2d',
]
class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
@@ -51,11 +62,13 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
use_bias: bool = False,
**kwargs
):
self.can_merge_in = True
"""if alpha == 0 or None, alpha is rank (no scaling)."""
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
self.lora_name = lora_name
self.scalar = torch.tensor(1.0)
self.orig_module_ref = weakref.ref(org_module)
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
@@ -111,10 +124,14 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
NUM_OF_BLOCKS = 12 # フルモデル相当でのup,downの層の数
UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel"]
# UNET_TARGET_REPLACE_MODULE = ["Transformer2DModel", "ResnetBlock2D"]
UNET_TARGET_REPLACE_MODULE = ["UNet2DConditionModel"]
# UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["ResnetBlock2D", "Downsample2D", "Upsample2D"]
UNET_TARGET_REPLACE_MODULE_CONV2D_3X3 = ["UNet2DConditionModel"]
TEXT_ENCODER_TARGET_REPLACE_MODULE = ["CLIPAttention", "CLIPMLP"]
LORA_PREFIX_UNET = "lora_unet"
PEFT_PREFIX_UNET = "unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
# SDXL: must starts with LORA_PREFIX_TEXT_ENCODER
@@ -147,12 +164,26 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
train_unet: Optional[bool] = True,
is_sdxl=False,
is_v2=False,
is_v3=False,
is_pixart: bool = False,
is_auraflow: bool = False,
is_flux: bool = False,
is_lumina2: bool = False,
use_bias: bool = False,
is_lorm: bool = False,
ignore_if_contains = None,
only_if_contains = None,
parameter_threshold: float = 0.0,
attn_only: bool = False,
target_lin_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE,
target_conv_modules=LoRANetwork.UNET_TARGET_REPLACE_MODULE_CONV2D_3X3,
network_type: str = "lora",
full_train_in_out: bool = False,
transformer_only: bool = False,
peft_format: bool = False,
is_assistant_adapter: bool = False,
is_transformer: bool = False,
base_model: 'StableDiffusion' = None,
**kwargs
) -> None:
"""
@@ -177,6 +208,13 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
if ignore_if_contains is None:
ignore_if_contains = []
self.ignore_if_contains = ignore_if_contains
self.transformer_only = transformer_only
self.base_model_ref = None
if base_model is not None:
self.base_model_ref = weakref.ref(base_model)
self.only_if_contains: Union[List, None] = only_if_contains
self.lora_dim = lora_dim
self.alpha = alpha
self.conv_lora_dim = conv_lora_dim
@@ -192,6 +230,39 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
self.multiplier = multiplier
self.is_sdxl = is_sdxl
self.is_v2 = is_v2
self.is_v3 = is_v3
self.is_pixart = is_pixart
self.is_auraflow = is_auraflow
self.is_flux = is_flux
self.is_lumina2 = is_lumina2
self.network_type = network_type
self.is_assistant_adapter = is_assistant_adapter
if self.network_type.lower() == "dora":
self.module_class = DoRAModule
module_class = DoRAModule
elif self.network_type.lower() == "lokr":
self.module_class = LokrModule
module_class = LokrModule
self.network_config: NetworkConfig = kwargs.get("network_config", None)
self.peft_format = peft_format
self.is_transformer = is_transformer
# always do peft for flux only for now
if self.is_flux or self.is_v3 or self.is_lumina2 or is_transformer:
# don't do peft format for lokr
if self.network_type.lower() != "lokr":
self.peft_format = True
if self.peft_format:
# no alpha for peft
self.alpha = self.lora_dim
alpha = self.alpha
self.conv_alpha = self.conv_lora_dim
conv_alpha = self.conv_alpha
self.full_train_in_out = full_train_in_out
if modules_dim is not None:
print(f"create LoRA network from weights")
@@ -219,8 +290,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
root_module: torch.nn.Module,
target_replace_modules: List[torch.nn.Module],
) -> List[LoRAModule]:
unet_prefix = self.LORA_PREFIX_UNET
if self.peft_format:
unet_prefix = self.PEFT_PREFIX_UNET
if is_pixart or is_v3 or is_auraflow or is_flux or is_lumina2 or self.is_transformer:
unet_prefix = f"lora_transformer"
if self.peft_format:
unet_prefix = "transformer"
prefix = (
self.LORA_PREFIX_UNET
unet_prefix
if is_unet
else (
self.LORA_PREFIX_TEXT_ENCODER
@@ -230,6 +309,8 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
)
loras = []
skipped = []
attached_modules = []
lora_shape_dict = {}
for name, module in root_module.named_modules():
if module.__class__.__name__ in target_replace_modules:
for child_name, child_module in module.named_modules():
@@ -237,17 +318,55 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
is_conv2d = child_module.__class__.__name__ in CONV_MODULES
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
lora_name = [prefix, name, child_name]
# filter out blank
lora_name = [x for x in lora_name if x and x != ""]
lora_name = ".".join(lora_name)
# if it doesnt have a name, it wil have two dots
lora_name.replace("..", ".")
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
if count_parameters(child_module) < parameter_threshold:
skip = True
if self.transformer_only and self.is_pixart and is_unet:
if "transformer_blocks" not in lora_name:
skip = True
if self.transformer_only and self.is_flux and is_unet:
if "transformer_blocks" not in lora_name:
skip = True
if self.transformer_only and self.is_lumina2 and is_unet:
if "layers$$" not in lora_name and "noise_refiner$$" not in lora_name and "context_refiner$$" not in lora_name:
skip = True
if self.transformer_only and self.is_v3 and is_unet:
if "transformer_blocks" not in lora_name:
skip = True
# handle custom models
if self.transformer_only and is_unet and hasattr(root_module, 'transformer_blocks'):
if "transformer_blocks" not in lora_name:
skip = True
if self.transformer_only and is_unet and hasattr(root_module, 'blocks'):
if "blocks" not in lora_name:
skip = True
if (is_linear or is_conv2d) and not skip:
lora_name = prefix + "." + name + "." + child_name
lora_name = lora_name.replace(".", "_")
if self.only_if_contains is not None:
if not any([word in clean_name for word in self.only_if_contains]) and not any([word in lora_name for word in self.only_if_contains]):
continue
dim = None
alpha = None
@@ -281,6 +400,11 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
self.conv_lora_dim is not None or conv_block_dims is not None):
skipped.append(lora_name)
continue
module_kwargs = {}
if self.network_type.lower() == "lokr":
module_kwargs["factor"] = self.network_config.lokr_factor
lora = module_class(
lora_name,
@@ -294,8 +418,16 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
network=self,
parent=module,
use_bias=use_bias,
**module_kwargs
)
loras.append(lora)
if self.network_type.lower() == "lokr":
try:
lora_shape_dict[lora_name] = [list(lora.lokr_w1.weight.shape), list(lora.lokr_w2.weight.shape)]
except:
pass
else:
lora_shape_dict[lora_name] = [list(lora.lora_down.weight.shape), list(lora.lora_up.weight.shape)]
return loras, skipped
text_encoders = text_encoder if type(text_encoder) == list else [text_encoder]
@@ -317,8 +449,12 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
index = None
print(f"create LoRA for Text Encoder:")
text_encoder_loras, skipped = create_modules(False, index, text_encoder,
LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE)
replace_modules = LoRANetwork.TEXT_ENCODER_TARGET_REPLACE_MODULE
if self.is_pixart:
replace_modules = ["T5EncoderModel"]
text_encoder_loras, skipped = create_modules(False, index, text_encoder, replace_modules)
self.text_encoder_loras.extend(text_encoder_loras)
skipped_te += skipped
print(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
@@ -328,6 +464,21 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
if modules_dim is not None or self.conv_lora_dim is not None or conv_block_dims is not None:
target_modules += target_conv_modules
if is_v3:
target_modules = ["SD3Transformer2DModel"]
if is_pixart:
target_modules = ["PixArtTransformer2DModel"]
if is_auraflow:
target_modules = ["AuraFlowTransformer2DModel"]
if is_flux:
target_modules = ["FluxTransformer2DModel"]
if is_lumina2:
target_modules = ["Lumina2Transformer2DModel"]
if train_unet:
self.unet_loras, skipped_un = create_modules(True, None, unet, target_modules)
else:
@@ -353,3 +504,49 @@ class LoRASpecialNetwork(ToolkitNetworkMixin, LoRANetwork):
for lora in self.text_encoder_loras + self.unet_loras:
assert lora.lora_name not in names, f"duplicated lora name: {lora.lora_name}"
names.add(lora.lora_name)
if self.full_train_in_out:
print("full train in out")
# we are going to retrain the main in out layers for VAE change usually
if self.is_pixart:
transformer: PixArtTransformer2DModel = unet
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
transformer.pos_embed = self.transformer_pos_embed
transformer.proj_out = self.transformer_proj_out
elif self.is_auraflow:
transformer: AuraFlowTransformer2DModel = unet
self.transformer_pos_embed = copy.deepcopy(transformer.pos_embed)
self.transformer_proj_out = copy.deepcopy(transformer.proj_out)
transformer.pos_embed = self.transformer_pos_embed
transformer.proj_out = self.transformer_proj_out
else:
unet: UNet2DConditionModel = unet
unet_conv_in: torch.nn.Conv2d = unet.conv_in
unet_conv_out: torch.nn.Conv2d = unet.conv_out
# clone these and replace their forwards with ours
self.unet_conv_in = copy.deepcopy(unet_conv_in)
self.unet_conv_out = copy.deepcopy(unet_conv_out)
unet.conv_in = self.unet_conv_in
unet.conv_out = self.unet_conv_out
def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr):
# call Lora prepare_optimizer_params
all_params = super().prepare_optimizer_params(text_encoder_lr, unet_lr, default_lr)
if self.full_train_in_out:
if self.is_pixart or self.is_auraflow or self.is_flux:
all_params.append({"lr": unet_lr, "params": list(self.transformer_pos_embed.parameters())})
all_params.append({"lr": unet_lr, "params": list(self.transformer_proj_out.parameters())})
else:
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_in.parameters())})
all_params.append({"lr": unet_lr, "params": list(self.unet_conv_out.parameters())})
return all_params

View File

@@ -354,7 +354,8 @@ def convert_diffusers_unet_to_lorm(
elif child_module.__class__.__name__ in LINEAR_MODULES:
if count_parameters(child_module) > parameter_threshold:
dtype = child_module.weight.dtype
# dtype = child_module.weight.dtype
dtype = torch.float32
# extract and convert
down_weight, up_weight, lora_dim, diff = extract_linear(
weight=child_module.weight.clone().detach().float(),

View File

@@ -23,6 +23,8 @@ def get_meta_for_safetensors(meta: OrderedDict, name=None, add_software_info=Tru
# if not float, int, bool, or str, convert to json string
if not isinstance(value, str):
save_meta[key] = json.dumps(value)
# add the pt format
save_meta["format"] = "pt"
return save_meta

146
toolkit/models/DoRA.py Normal file
View File

@@ -0,0 +1,146 @@
#based off https://github.com/catid/dora/blob/main/dora.py
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import TYPE_CHECKING, Union, List
from optimum.quanto import QBytesTensor, QTensor
from toolkit.network_mixins import ToolkitModuleMixin, ExtractableModuleMixin
if TYPE_CHECKING:
from toolkit.lora_special import LoRASpecialNetwork
# diffusers specific stuff
LINEAR_MODULES = [
'Linear',
'LoRACompatibleLinear'
# 'GroupNorm',
]
CONV_MODULES = [
'Conv2d',
'LoRACompatibleConv'
]
def transpose(weight, fan_in_fan_out):
if not fan_in_fan_out:
return weight
if isinstance(weight, torch.nn.Parameter):
return torch.nn.Parameter(weight.T)
return weight.T
class DoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module):
# def __init__(self, d_in, d_out, rank=4, weight=None, bias=None):
def __init__(
self,
lora_name,
org_module: torch.nn.Module,
multiplier=1.0,
lora_dim=4,
alpha=1,
dropout=None,
rank_dropout=None,
module_dropout=None,
network: 'LoRASpecialNetwork' = None,
use_bias: bool = False,
**kwargs
):
self.can_merge_in = False
"""if alpha == 0 or None, alpha is rank (no scaling)."""
ToolkitModuleMixin.__init__(self, network=network)
torch.nn.Module.__init__(self)
self.lora_name = lora_name
self.scalar = torch.tensor(1.0)
self.lora_dim = lora_dim
if org_module.__class__.__name__ in CONV_MODULES:
raise NotImplementedError("Convolutional layers are not supported yet")
if type(alpha) == torch.Tensor:
alpha = alpha.detach().float().numpy() # without casting, bf16 causes error
alpha = self.lora_dim if alpha is None or alpha == 0 else alpha
self.scale = alpha / self.lora_dim
# self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える eng: treat as constant
self.multiplier: Union[float, List[float]] = multiplier
# wrap the original module so it doesn't get weights updated
self.org_module = [org_module]
self.dropout = dropout
self.rank_dropout = rank_dropout
self.module_dropout = module_dropout
self.is_checkpointing = False
d_out = org_module.out_features
d_in = org_module.in_features
std_dev = 1 / torch.sqrt(torch.tensor(self.lora_dim).float())
# self.lora_up = nn.Parameter(torch.randn(d_out, self.lora_dim) * std_dev) # lora_A
# self.lora_down = nn.Parameter(torch.zeros(self.lora_dim, d_in)) # lora_B
self.lora_up = nn.Linear(self.lora_dim, d_out, bias=False) # lora_B
# self.lora_up.weight.data = torch.randn_like(self.lora_up.weight.data) * std_dev
self.lora_up.weight.data = torch.zeros_like(self.lora_up.weight.data)
# self.lora_A[adapter_name] = nn.Linear(self.in_features, r, bias=False)
# self.lora_B[adapter_name] = nn.Linear(r, self.out_features, bias=False)
self.lora_down = nn.Linear(d_in, self.lora_dim, bias=False) # lora_A
# self.lora_down.weight.data = torch.zeros_like(self.lora_down.weight.data)
self.lora_down.weight.data = torch.randn_like(self.lora_down.weight.data) * std_dev
# m = Magnitude column-wise across output dimension
weight = self.get_orig_weight()
weight = weight.to(self.lora_up.weight.device, dtype=self.lora_up.weight.dtype)
lora_weight = self.lora_up.weight @ self.lora_down.weight
weight_norm = self._get_weight_norm(weight, lora_weight)
self.magnitude = nn.Parameter(weight_norm.detach().clone(), requires_grad=True)
def apply_to(self):
self.org_forward = self.org_module[0].forward
self.org_module[0].forward = self.forward
# del self.org_module
def get_orig_weight(self):
weight = self.org_module[0].weight
if isinstance(weight, QTensor) or isinstance(weight, QBytesTensor):
return weight.dequantize().data.detach()
else:
return weight.data.detach()
def get_orig_bias(self):
if hasattr(self.org_module[0], 'bias') and self.org_module[0].bias is not None:
return self.org_module[0].bias.data.detach()
return None
# def dora_forward(self, x, *args, **kwargs):
# lora = torch.matmul(self.lora_A, self.lora_B)
# adapted = self.get_orig_weight() + lora
# column_norm = adapted.norm(p=2, dim=0, keepdim=True)
# norm_adapted = adapted / column_norm
# calc_weights = self.magnitude * norm_adapted
# return F.linear(x, calc_weights, self.get_orig_bias())
def _get_weight_norm(self, weight, scaled_lora_weight) -> torch.Tensor:
# calculate L2 norm of weight matrix, column-wise
weight = weight + scaled_lora_weight.to(weight.device)
weight_norm = torch.linalg.norm(weight, dim=1)
return weight_norm
def apply_dora(self, x, scaled_lora_weight):
# ref https://github.com/huggingface/peft/blob/1e6d1d73a0850223b0916052fd8d2382a90eae5a/src/peft/tuners/lora/layer.py#L192
# lora weight is already scaled
# magnitude = self.lora_magnitude_vector[active_adapter]
weight = self.get_orig_weight()
weight = weight.to(scaled_lora_weight.device, dtype=scaled_lora_weight.dtype)
weight_norm = self._get_weight_norm(weight, scaled_lora_weight)
# see section 4.3 of DoRA (https://arxiv.org/abs/2402.09353)
# "[...] we suggest treating ||V +∆V ||_c in
# Eq. (5) as a constant, thereby detaching it from the gradient
# graph. This means that while ||V + ∆V ||_c dynamically
# reflects the updates of ∆V , it won’t receive any gradient
# during backpropagation"
weight_norm = weight_norm.detach()
dora_weight = transpose(weight + scaled_lora_weight, False)
return (self.magnitude / weight_norm - 1).view(1, -1) * F.linear(x.to(dora_weight.dtype), dora_weight)

View File

@@ -0,0 +1,267 @@
import math
import weakref
import torch
import torch.nn as nn
from typing import TYPE_CHECKING, List, Dict, Any
from toolkit.models.clip_fusion import ZipperBlock
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
import sys
from toolkit.paths import REPOS_ROOT
sys.path.append(REPOS_ROOT)
from ipadapter.ip_adapter.resampler import Resampler
from collections import OrderedDict
if TYPE_CHECKING:
from toolkit.lora_special import LoRAModule
from toolkit.stable_diffusion_model import StableDiffusion
class TransformerBlock(nn.Module):
def __init__(self, d_model, nhead, dim_feedforward):
super().__init__()
self.self_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
self.cross_attn = nn.MultiheadAttention(d_model, nhead, batch_first=True)
self.feed_forward = nn.Sequential(
nn.Linear(d_model, dim_feedforward),
nn.ReLU(),
nn.Linear(dim_feedforward, d_model)
)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
def forward(self, x, cross_attn_input):
# Self-attention
attn_output, _ = self.self_attn(x, x, x)
x = self.norm1(x + attn_output)
# Cross-attention
cross_attn_output, _ = self.cross_attn(x, cross_attn_input, cross_attn_input)
x = self.norm2(x + cross_attn_output)
# Feed-forward
ff_output = self.feed_forward(x)
x = self.norm3(x + ff_output)
return x
class InstantLoRAMidModule(torch.nn.Module):
def __init__(
self,
index: int,
lora_module: 'LoRAModule',
instant_lora_module: 'InstantLoRAModule',
up_shape: list = None,
down_shape: list = None,
):
super(InstantLoRAMidModule, self).__init__()
self.up_shape = up_shape
self.down_shape = down_shape
self.index = index
self.lora_module_ref = weakref.ref(lora_module)
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
self.embed = None
def down_forward(self, x, *args, **kwargs):
# get the embed
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
down_size = math.prod(self.down_shape)
down_weight = self.embed[:, :down_size]
batch_size = x.shape[0]
# unconditional
if down_weight.shape[0] * 2 == batch_size:
down_weight = torch.cat([down_weight] * 2, dim=0)
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
x_chunks = torch.chunk(x, batch_size, dim=0)
x_out = []
for i in range(batch_size):
weight_chunk = weight_chunks[i]
x_chunk = x_chunks[i]
# reshape
weight_chunk = weight_chunk.view(self.down_shape)
# check if is conv or linear
if len(weight_chunk.shape) == 4:
padding = 0
if weight_chunk.shape[-1] == 3:
padding = 1
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
else:
# run a simple linear layer with the down weight
x_chunk = x_chunk @ weight_chunk.T
x_out.append(x_chunk)
x = torch.cat(x_out, dim=0)
return x
def up_forward(self, x, *args, **kwargs):
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
up_size = math.prod(self.up_shape)
up_weight = self.embed[:, -up_size:]
batch_size = x.shape[0]
# unconditional
if up_weight.shape[0] * 2 == batch_size:
up_weight = torch.cat([up_weight] * 2, dim=0)
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
x_chunks = torch.chunk(x, batch_size, dim=0)
x_out = []
for i in range(batch_size):
weight_chunk = weight_chunks[i]
x_chunk = x_chunks[i]
# reshape
weight_chunk = weight_chunk.view(self.up_shape)
# check if is conv or linear
if len(weight_chunk.shape) == 4:
padding = 0
if weight_chunk.shape[-1] == 3:
padding = 1
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
else:
# run a simple linear layer with the down weight
x_chunk = x_chunk @ weight_chunk.T
x_out.append(x_chunk)
x = torch.cat(x_out, dim=0)
return x
# Initialize the network
# num_blocks = 8
# d_model = 1024 # Adjust as needed
# nhead = 16 # Adjust as needed
# dim_feedforward = 4096 # Adjust as needed
# latent_dim = 1695744
class LoRAFormer(torch.nn.Module):
def __init__(
self,
num_blocks,
d_model=1024,
nhead=16,
dim_feedforward=4096,
sd: 'StableDiffusion'=None,
):
super(LoRAFormer, self).__init__()
# self.linear = torch.nn.Linear(2, 1)
self.sd_ref = weakref.ref(sd)
self.dim = sd.network.lora_dim
# stores the projection vector. Grabbed by modules
self.img_embeds: List[torch.Tensor] = None
# disable merging in. It is slower on inference
self.sd_ref().network.can_merge_in = False
self.ilora_modules = torch.nn.ModuleList()
lora_modules = self.sd_ref().network.get_all_modules()
output_size = 0
self.embed_lengths = []
self.weight_mapping = []
for idx, lora_module in enumerate(lora_modules):
module_dict = lora_module.state_dict()
down_shape = list(module_dict['lora_down.weight'].shape)
up_shape = list(module_dict['lora_up.weight'].shape)
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
module_size = math.prod(down_shape) + math.prod(up_shape)
output_size += module_size
self.embed_lengths.append(module_size)
# add a new mid module that will take the original forward and add a vector to it
# this will be used to add the vector to the original forward
instant_module = InstantLoRAMidModule(
idx,
lora_module,
self,
up_shape=up_shape,
down_shape=down_shape
)
self.ilora_modules.append(instant_module)
# replace the LoRA forwards
lora_module.lora_down.forward = instant_module.down_forward
lora_module.lora_up.forward = instant_module.up_forward
self.output_size = output_size
self.latent = nn.Parameter(torch.randn(1, output_size))
self.latent_proj = nn.Linear(output_size, d_model)
self.blocks = nn.ModuleList([
TransformerBlock(d_model, nhead, dim_feedforward)
for _ in range(num_blocks)
])
self.final_proj = nn.Linear(d_model, output_size)
self.migrate_weight_mapping()
def migrate_weight_mapping(self):
return
# # changes the names of the modules to common ones
# keymap = self.sd_ref().network.get_keymap()
# save_keymap = {}
# if keymap is not None:
# for ldm_key, diffusers_key in keymap.items():
# # invert them
# save_keymap[diffusers_key] = ldm_key
#
# new_keymap = {}
# for key, value in self.weight_mapping:
# if key in save_keymap:
# new_keymap[save_keymap[key]] = value
# else:
# print(f"Key {key} not found in keymap")
# new_keymap[key] = value
# self.weight_mapping = new_keymap
# else:
# print("No keymap found. Using default names")
# return
def forward(self, img_embeds):
# expand token rank if only rank 2
if len(img_embeds.shape) == 2:
img_embeds = img_embeds.unsqueeze(1)
# resample the image embeddings
img_embeds = self.resampler(img_embeds)
img_embeds = self.proj_module(img_embeds)
if len(img_embeds.shape) == 3:
# merge the heads
img_embeds = img_embeds.mean(dim=1)
self.img_embeds = []
# get all the slices
start = 0
for length in self.embed_lengths:
self.img_embeds.append(img_embeds[:, start:start+length])
start += length
def get_additional_save_metadata(self) -> Dict[str, Any]:
# save the weight mapping
return {
"weight_mapping": self.weight_mapping,
"num_heads": self.num_heads,
"vision_hidden_size": self.vision_hidden_size,
"head_dim": self.head_dim,
"vision_tokens": self.vision_tokens,
"output_size": self.output_size,
}

127
toolkit/models/auraflow.py Normal file
View File

@@ -0,0 +1,127 @@
import math
from functools import partial
from torch import nn
import torch
class AuraFlowPatchEmbed(nn.Module):
def __init__(
self,
height=224,
width=224,
patch_size=16,
in_channels=3,
embed_dim=768,
pos_embed_max_size=None,
):
super().__init__()
self.num_patches = (height // patch_size) * (width // patch_size)
self.pos_embed_max_size = pos_embed_max_size
self.proj = nn.Linear(patch_size * patch_size * in_channels, embed_dim)
self.pos_embed = nn.Parameter(torch.randn(1, pos_embed_max_size, embed_dim) * 0.1)
self.patch_size = patch_size
self.height, self.width = height // patch_size, width // patch_size
self.base_size = height // patch_size
def forward(self, latent):
batch_size, num_channels, height, width = latent.size()
latent = latent.view(
batch_size,
num_channels,
height // self.patch_size,
self.patch_size,
width // self.patch_size,
self.patch_size,
)
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
latent = self.proj(latent)
try:
return latent + self.pos_embed
except RuntimeError:
raise RuntimeError(
f"Positional embeddings are too small for the number of patches. "
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
)
# comfy
# def apply_pos_embeds(self, x, h, w):
# h = (h + 1) // self.patch_size
# w = (w + 1) // self.patch_size
# max_dim = max(h, w)
#
# cur_dim = self.h_max
# pos_encoding = self.positional_encoding.reshape(1, cur_dim, cur_dim, -1).to(device=x.device, dtype=x.dtype)
#
# if max_dim > cur_dim:
# pos_encoding = F.interpolate(pos_encoding.movedim(-1, 1), (max_dim, max_dim), mode="bilinear").movedim(1,
# -1)
# cur_dim = max_dim
#
# from_h = (cur_dim - h) // 2
# from_w = (cur_dim - w) // 2
# pos_encoding = pos_encoding[:, from_h:from_h + h, from_w:from_w + w]
# return x + pos_encoding.reshape(1, -1, self.positional_encoding.shape[-1])
# def patchify(self, x):
# B, C, H, W = x.size()
# pad_h = (self.patch_size - H % self.patch_size) % self.patch_size
# pad_w = (self.patch_size - W % self.patch_size) % self.patch_size
#
# x = torch.nn.functional.pad(x, (0, pad_w, 0, pad_h), mode='reflect')
# x = x.view(
# B,
# C,
# (H + 1) // self.patch_size,
# self.patch_size,
# (W + 1) // self.patch_size,
# self.patch_size,
# )
# x = x.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
# return x
def patch_auraflow_pos_embed(pos_embed):
# we need to hijack the forward and replace with a custom one. Self is the model
def new_forward(self, latent):
batch_size, num_channels, height, width = latent.size()
# add padding to the latent to make it match pos_embed
latent_size = height * width * num_channels / 16 # todo check where 16 comes from?
pos_embed_size = self.pos_embed.shape[1]
if latent_size < pos_embed_size:
total_padding = int(pos_embed_size - math.floor(latent_size))
total_padding = total_padding // 16
pad_height = total_padding // 2
pad_width = total_padding - pad_height
# mirror padding on the right side
padding = (0, pad_width, 0, pad_height)
latent = torch.nn.functional.pad(latent, padding, mode='reflect')
elif latent_size > pos_embed_size:
amount_to_remove = latent_size - pos_embed_size
latent = latent[:, :, :-amount_to_remove]
batch_size, num_channels, height, width = latent.size()
latent = latent.view(
batch_size,
num_channels,
height // self.patch_size,
self.patch_size,
width // self.patch_size,
self.patch_size,
)
latent = latent.permute(0, 2, 4, 1, 3, 5).flatten(-3).flatten(1, 2)
latent = self.proj(latent)
try:
return latent + self.pos_embed
except RuntimeError:
raise RuntimeError(
f"Positional embeddings are too small for the number of patches. "
f"Please increase `pos_embed_max_size` to at least {self.num_patches}."
)
pos_embed.forward = partial(new_forward, pos_embed)

1433
toolkit/models/base_model.py Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,162 @@
import torch
import torch.nn as nn
from toolkit.models.zipper_resampler import ContextualAlphaMask
# Conv1d MLP
# MLP that can alternately be used as a conv1d on dim 1
class MLPC(nn.Module):
def __init__(
self,
in_dim,
out_dim,
hidden_dim,
do_conv=False,
use_residual=True
):
super().__init__()
self.do_conv = do_conv
if use_residual:
assert in_dim == out_dim
# dont normalize if using conv
if not do_conv:
self.layernorm = nn.LayerNorm(in_dim)
if do_conv:
self.fc1 = nn.Conv1d(in_dim, hidden_dim, 1)
self.fc2 = nn.Conv1d(hidden_dim, out_dim, 1)
else:
self.fc1 = nn.Linear(in_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, out_dim)
self.use_residual = use_residual
self.act_fn = nn.GELU()
def forward(self, x):
residual = x
if not self.do_conv:
x = self.layernorm(x)
x = self.fc1(x)
x = self.act_fn(x)
x = self.fc2(x)
if self.use_residual:
x = x + residual
return x
class ZipperBlock(nn.Module):
def __init__(
self,
in_size,
in_tokens,
out_size,
out_tokens,
hidden_size,
hidden_tokens,
):
super().__init__()
self.in_size = in_size
self.in_tokens = in_tokens
self.out_size = out_size
self.out_tokens = out_tokens
self.hidden_size = hidden_size
self.hidden_tokens = hidden_tokens
# permute to (batch_size, out_size, in_tokens)
self.zip_token = MLPC(
in_dim=self.in_tokens,
out_dim=self.out_tokens,
hidden_dim=self.hidden_tokens,
do_conv=True, # no need to permute
use_residual=False
)
# permute to (batch_size, out_tokens, out_size)
# in shpae: (batch_size, in_tokens, in_size)
self.zip_size = MLPC(
in_dim=self.in_size,
out_dim=self.out_size,
hidden_dim=self.hidden_size,
use_residual=False
)
def forward(self, x):
x = self.zip_token(x)
x = self.zip_size(x)
return x
# CLIPFusionModule
# Fuses any size of vision and text embeddings into a single embedding.
# remaps tokens and vectors.
class CLIPFusionModule(nn.Module):
def __init__(
self,
text_hidden_size: int = 768,
text_tokens: int = 77,
vision_hidden_size: int = 1024,
vision_tokens: int = 257,
num_blocks: int = 1,
):
super(CLIPFusionModule, self).__init__()
self.text_hidden_size = text_hidden_size
self.text_tokens = text_tokens
self.vision_hidden_size = vision_hidden_size
self.vision_tokens = vision_tokens
self.resampler = ZipperBlock(
in_size=self.vision_hidden_size,
in_tokens=self.vision_tokens,
out_size=self.text_hidden_size,
out_tokens=self.text_tokens,
hidden_size=self.vision_hidden_size * 2,
hidden_tokens=self.vision_tokens * 2
)
self.zipper_blocks = torch.nn.ModuleList([
ZipperBlock(
in_size=self.text_hidden_size * 2,
in_tokens=self.text_tokens,
out_size=self.text_hidden_size,
out_tokens=self.text_tokens,
hidden_size=self.text_hidden_size * 2,
hidden_tokens=self.text_tokens * 2
) for i in range(num_blocks)
])
self.ctx_alpha = ContextualAlphaMask(
dim=self.text_hidden_size,
)
self.alpha = nn.Parameter(torch.zeros([text_tokens]) + 0.01)
def forward(self, text_embeds, vision_embeds):
# text_embeds = (batch_size, 77, 768)
# vision_embeds = (batch_size, 257, 1024)
# output = (batch_size, 77, 768)
vision_embeds = self.resampler(vision_embeds)
x = vision_embeds
for i, block in enumerate(self.zipper_blocks):
res = x
x = torch.cat([text_embeds, x], dim=-1)
x = block(x)
x = x + res
# alpha mask
ctx_alpha = self.ctx_alpha(text_embeds)
# reshape alpha to (1, 77, 1)
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
x = ctx_alpha * x * alpha
x = x + text_embeds
return x

View File

@@ -0,0 +1,123 @@
import torch
import torch.nn as nn
class UpsampleBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.conv_in = nn.Sequential(
nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1),
nn.GELU()
)
self.conv_up = nn.Sequential(
nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2, stride=2),
nn.GELU()
)
self.conv_out = nn.Sequential(
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
)
def forward(self, x):
x = self.conv_in(x)
x = self.conv_up(x)
x = self.conv_out(x)
return x
class CLIPImagePreProcessor(nn.Module):
def __init__(
self,
input_size=896,
clip_input_size=224,
downscale_factor: int = 16,
):
super().__init__()
# make sure they are evenly divisible
assert input_size % clip_input_size == 0
in_channels = 3
self.input_size = input_size
self.clip_input_size = clip_input_size
self.downscale_factor = downscale_factor
subpixel_channels = in_channels * downscale_factor ** 2 # 3 * 16 ** 2 = 768
channels = subpixel_channels
upscale_factor = downscale_factor / int((input_size / clip_input_size)) # 16 / (896 / 224) = 4
num_upsample_blocks = int(upscale_factor // 2) # 4 // 2 = 2
# make the residual down up blocks
self.upsample_blocks = nn.ModuleList()
self.subpixel_blocks = nn.ModuleList()
current_channels = channels
current_downscale = downscale_factor
for _ in range(num_upsample_blocks):
# determine the reshuffled channel count for this dimension
output_downscale = current_downscale // 2
out_channels = in_channels * output_downscale ** 2
# out_channels = current_channels // 2
self.upsample_blocks.append(UpsampleBlock(current_channels, out_channels))
current_channels = out_channels
current_downscale = output_downscale
self.subpixel_blocks.append(nn.PixelUnshuffle(current_downscale))
# (bs, 768, 56, 56) -> (bs, 192, 112, 112)
# (bs, 192, 112, 112) -> (bs, 48, 224, 224)
self.conv_out = nn.Conv2d(
current_channels,
out_channels=3,
kernel_size=3,
padding=1
) # (bs, 48, 224, 224) -> (bs, 3, 224, 224)
# do a pooling layer to downscale the input to 1/3 of the size
# (bs, 3, 896, 896) -> (bs, 3, 224, 224)
kernel_size = input_size // clip_input_size
self.res_down = nn.AvgPool2d(
kernel_size=kernel_size,
stride=kernel_size
) # (bs, 3, 896, 896) -> (bs, 3, 224, 224)
# make a blending for output residual with near 0 weight
self.res_blend = nn.Parameter(torch.tensor(0.001)) # (bs, 3, 224, 224) -> (bs, 3, 224, 224)
self.unshuffle = nn.PixelUnshuffle(downscale_factor) # (bs, 3, 896, 896) -> (bs, 768, 56, 56)
self.conv_in = nn.Sequential(
nn.Conv2d(
subpixel_channels,
channels,
kernel_size=3,
padding=1
),
nn.GELU()
) # (bs, 768, 56, 56) -> (bs, 768, 56, 56)
# make 2 deep blocks
def forward(self, x):
inputs = x
# resize to input_size x input_size
x = nn.functional.interpolate(x, size=(self.input_size, self.input_size), mode='bicubic')
res = self.res_down(inputs)
x = self.unshuffle(x)
x = self.conv_in(x)
for up, subpixel in zip(self.upsample_blocks, self.subpixel_blocks):
x = up(x)
block_res = subpixel(inputs)
x = x + block_res
x = self.conv_out(x)
# blend residual
x = x * self.res_blend + res
return x

466
toolkit/models/cogview4.py Normal file
View File

@@ -0,0 +1,466 @@
# DONT USE THIS!. IT DOES NOT WORK YET!
# Will revisit this when they release more info on how it was trained.
import weakref
from diffusers import CogView4Pipeline
import torch
import yaml
from toolkit.basic import flush
from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.dequantize import patch_dequantization_on_save
from toolkit.models.base_model import BaseModel
from toolkit.prompt_utils import PromptEmbeds
import os
import copy
from toolkit.config_modules import ModelConfig, GenerateImageConfig, ModelArch
import torch
import diffusers
from diffusers import AutoencoderKL, CogView4Transformer2DModel, CogView4Pipeline
from optimum.quanto import freeze, qfloat8, QTensor, qint4
from toolkit.util.quantize import quantize, get_qtype
from transformers import GlmModel, AutoTokenizer
from diffusers import FlowMatchEulerDiscreteScheduler
from typing import TYPE_CHECKING
from toolkit.accelerator import unwrap_model
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
if TYPE_CHECKING:
from toolkit.lora_special import LoRASpecialNetwork
# remove this after a bug is fixed in diffusers code. This is a workaround.
class FakeModel:
def __init__(self, model):
self.model_ref = weakref.ref(model)
pass
@property
def device(self):
return self.model_ref().device
scheduler_config = {
"base_image_seq_len": 256,
"base_shift": 0.25,
"invert_sigmas": False,
"max_image_seq_len": 4096,
"max_shift": 0.75,
"num_train_timesteps": 1000,
"shift": 1.0,
"shift_terminal": None,
"time_shift_type": "linear",
"use_beta_sigmas": False,
"use_dynamic_shifting": True,
"use_exponential_sigmas": False,
"use_karras_sigmas": False
}
class CogView4(BaseModel):
def __init__(
self,
device,
model_config: ModelConfig,
dtype='bf16',
custom_pipeline=None,
noise_scheduler=None,
**kwargs
):
super().__init__(device, model_config, dtype,
custom_pipeline, noise_scheduler, **kwargs)
self.is_flow_matching = True
self.is_transformer = True
self.target_lora_modules = ['CogView4Transformer2DModel']
# cache for holding noise
self.effective_noise = None
# static method to get the scheduler
@staticmethod
def get_train_scheduler():
scheduler = CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
return scheduler
def load_model(self):
dtype = self.torch_dtype
base_model_path = "THUDM/CogView4-6B"
model_path = self.model_config.name_or_path
self.print_and_status_update("Loading CogView4 model")
# base_model_path = "black-forest-labs/FLUX.1-schnell"
base_model_path = self.model_config.name_or_path_original
subfolder = 'transformer'
transformer_path = model_path
if os.path.exists(transformer_path):
subfolder = None
transformer_path = os.path.join(transformer_path, 'transformer')
# check if the path is a full checkpoint.
te_folder_path = os.path.join(model_path, 'text_encoder')
# if we have the te, this folder is a full checkpoint, use it as the base
if os.path.exists(te_folder_path):
base_model_path = model_path
self.print_and_status_update("Loading GlmModel")
tokenizer = AutoTokenizer.from_pretrained(
base_model_path, subfolder="tokenizer", torch_dtype=dtype)
text_encoder = GlmModel.from_pretrained(
base_model_path, subfolder="text_encoder", torch_dtype=dtype)
text_encoder.to(self.device_torch, dtype=dtype)
flush()
if self.model_config.quantize_te:
self.print_and_status_update("Quantizing GlmModel")
quantize(text_encoder, weights=get_qtype(self.model_config.qtype))
freeze(text_encoder)
flush()
# hack to fix diffusers bug workaround
text_encoder.model = FakeModel(text_encoder)
self.print_and_status_update("Loading transformer")
transformer = CogView4Transformer2DModel.from_pretrained(
transformer_path,
subfolder=subfolder,
torch_dtype=dtype,
)
if self.model_config.split_model_over_gpus:
raise ValueError(
"Splitting model over gpus is not supported for CogViewModels models")
transformer.to(self.quantize_device, dtype=dtype)
flush()
if self.model_config.assistant_lora_path is not None or self.model_config.inference_lora_path is not None:
raise ValueError(
"Assistant LoRA is not supported for CogViewModels models currently")
if self.model_config.lora_path is not None:
raise ValueError(
"Loading LoRA is not supported for CogViewModels models currently")
flush()
if self.model_config.quantize:
quantization_args = self.model_config.quantize_kwargs
if 'exclude' not in quantization_args:
quantization_args['exclude'] = []
if 'include' not in quantization_args:
quantization_args['include'] = []
# Be more specific with the include pattern to exactly match transformer blocks
quantization_args['include'] += ["transformer_blocks.*"]
# Exclude all LayerNorm layers within transformer blocks
quantization_args['exclude'] += [
"transformer_blocks.*.norm1",
"transformer_blocks.*.norm2",
"transformer_blocks.*.norm2_context",
"transformer_blocks.*.attn1.norm_q",
"transformer_blocks.*.attn1.norm_k"
]
# patch the state dict method
patch_dequantization_on_save(transformer)
quantization_type = get_qtype(self.model_config.qtype)
self.print_and_status_update("Quantizing transformer")
quantize(transformer, weights=quantization_type, **quantization_args)
freeze(transformer)
transformer.to(self.device_torch)
else:
transformer.to(self.device_torch, dtype=dtype)
flush()
scheduler = CogView4.get_train_scheduler()
self.print_and_status_update("Loading VAE")
vae = AutoencoderKL.from_pretrained(
base_model_path, subfolder="vae", torch_dtype=dtype)
flush()
self.print_and_status_update("Making pipe")
pipe: CogView4Pipeline = CogView4Pipeline(
scheduler=scheduler,
text_encoder=None,
tokenizer=tokenizer,
vae=vae,
transformer=None,
)
pipe.text_encoder = text_encoder
pipe.transformer = transformer
self.print_and_status_update("Preparing Model")
text_encoder = pipe.text_encoder
tokenizer = pipe.tokenizer
pipe.transformer = pipe.transformer.to(self.device_torch)
flush()
text_encoder.to(self.device_torch)
text_encoder.requires_grad_(False)
text_encoder.eval()
pipe.transformer = pipe.transformer.to(self.device_torch)
flush()
self.pipeline = pipe
self.model = transformer
self.vae = vae
self.text_encoder = text_encoder
self.tokenizer = tokenizer
def get_generation_pipeline(self):
scheduler = CogView4.get_train_scheduler()
pipeline = CogView4Pipeline(
vae=self.vae,
transformer=self.unet,
text_encoder=self.text_encoder,
tokenizer=self.tokenizer,
scheduler=scheduler,
)
return pipeline
def generate_single_image(
self,
pipeline: CogView4Pipeline,
gen_config: GenerateImageConfig,
conditional_embeds: PromptEmbeds,
unconditional_embeds: PromptEmbeds,
generator: torch.Generator,
extra: dict,
):
img = pipeline(
prompt_embeds=conditional_embeds.text_embeds.to(
self.device_torch, dtype=self.torch_dtype),
negative_prompt_embeds=unconditional_embeds.text_embeds.to(
self.device_torch, dtype=self.torch_dtype),
height=gen_config.height,
width=gen_config.width,
num_inference_steps=gen_config.num_inference_steps,
guidance_scale=gen_config.guidance_scale,
latents=gen_config.latents,
generator=generator,
**extra
).images[0]
return img
def get_noise_prediction(
self,
latent_model_input: torch.Tensor,
timestep: torch.Tensor, # 0 to 1000 scale
text_embeddings: PromptEmbeds,
**kwargs
):
# target_size = (height, width)
target_size = latent_model_input.shape[-2:]
# multiply by 8
target_size = (target_size[0] * 8, target_size[1] * 8)
crops_coords_top_left = torch.tensor(
[(0, 0)], dtype=self.torch_dtype, device=self.device_torch)
original_size = torch.tensor(
[target_size], dtype=self.torch_dtype, device=self.device_torch)
target_size = original_size.clone()
noise_pred_cond = self.model(
hidden_states=latent_model_input,
encoder_hidden_states=text_embeddings.text_embeds,
timestep=timestep,
original_size=original_size,
target_size=target_size,
crop_coords=crops_coords_top_left,
return_dict=False,
)[0]
return noise_pred_cond
def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
prompt_embeds, _ = self.pipeline.encode_prompt(
prompt,
do_classifier_free_guidance=False,
device=self.device_torch,
dtype=self.torch_dtype,
)
return PromptEmbeds(prompt_embeds)
def get_model_has_grad(self):
return self.model.proj_out.weight.requires_grad
def get_te_has_grad(self):
return self.text_encoder.layers[0].mlp.down_proj.weight.requires_grad
def save_model(self, output_path, meta, save_dtype):
# only save the unet
transformer: CogView4Transformer2DModel = unwrap_model(self.model)
transformer.save_pretrained(
save_directory=os.path.join(output_path, 'transformer'),
safe_serialization=True,
)
meta_path = os.path.join(output_path, 'aitk_meta.yaml')
with open(meta_path, 'w') as f:
yaml.dump(meta, f)
def get_loss_target(self, *args, **kwargs):
noise = kwargs.get('noise')
effective_noise = self.effective_noise
batch = kwargs.get('batch')
if batch is None:
raise ValueError("Batch is not provided")
if noise is None:
raise ValueError("Noise is not provided")
# return batch.latents
# return (batch.latents - noise).detach()
return (noise - batch.latents).detach()
# return (batch.latents).detach()
# return (effective_noise - batch.latents).detach()
def _get_low_res_latents(self, latents):
# todo prevent needing to do this and grab the tensor another way.
with torch.no_grad():
# Decode latents to image space
images = self.decode_latents(
latents, device=latents.device, dtype=latents.dtype)
# Downsample by a factor of 2 using bilinear interpolation
B, C, H, W = images.shape
low_res_images = torch.nn.functional.interpolate(
images,
size=(H // 2, W // 2),
mode="bilinear",
align_corners=False
)
# Upsample back to original resolution to match expected VAE input dimensions
upsampled_low_res_images = torch.nn.functional.interpolate(
low_res_images,
size=(H, W),
mode="bilinear",
align_corners=False
)
# Encode the low-resolution images back to latent space
low_res_latents = self.encode_images(
upsampled_low_res_images, device=latents.device, dtype=latents.dtype)
return low_res_latents
# def add_noise(
# self,
# original_samples: torch.FloatTensor,
# noise: torch.FloatTensor,
# timesteps: torch.IntTensor,
# **kwargs,
# ) -> torch.FloatTensor:
# relay_start_point = 500
# # Store original samples for loss calculation
# self.original_samples = original_samples
# # Prepare chunks for batch processing
# original_samples_chunks = torch.chunk(
# original_samples, original_samples.shape[0], dim=0)
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
# # Get the low res latents only if needed
# low_res_latents_chunks = None
# # Handle case where timesteps is a single value for all samples
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
# noisy_latents_chunks = []
# effective_noise_chunks = [] # Store the effective noise for each sample
# for idx in range(original_samples.shape[0]):
# t = timesteps_chunks[idx]
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
# # Flowmatching interpolation between original and noise
# if t > relay_start_point:
# # Standard flowmatching - direct linear interpolation
# noisy_latents = (1 - t_01) * original_samples_chunks[idx] + t_01 * noise_chunks[idx]
# effective_noise_chunks.append(noise_chunks[idx]) # Effective noise is just the noise
# else:
# # Relay flowmatching case - only compute low_res_latents if needed
# if low_res_latents_chunks is None:
# low_res_latents = self._get_low_res_latents(original_samples)
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
# # Calculate the relay ratio (0 to 1)
# t_ratio = t.float() / relay_start_point
# t_ratio = torch.clamp(t_ratio, 0.0, 1.0)
# # First blend between original and low-res based on t_ratio
# z0_t = (1 - t_ratio) * original_samples_chunks[idx] + t_ratio * low_res_latents_chunks[idx]
# added_lor_res_noise = z0_t - original_samples_chunks[idx]
# # Then apply flowmatching interpolation between this blended state and noise
# noisy_latents = (1 - t_01) * z0_t + t_01 * noise_chunks[idx]
# # For prediction target, we need to store the effective "source"
# effective_noise_chunks.append(noise_chunks[idx] + added_lor_res_noise)
# noisy_latents_chunks.append(noisy_latents)
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
# return noisy_latents
# def add_noise(
# self,
# original_samples: torch.FloatTensor,
# noise: torch.FloatTensor,
# timesteps: torch.IntTensor,
# **kwargs,
# ) -> torch.FloatTensor:
# relay_start_point = 500
# # Store original samples for loss calculation
# self.original_samples = original_samples
# # Prepare chunks for batch processing
# original_samples_chunks = torch.chunk(
# original_samples, original_samples.shape[0], dim=0)
# noise_chunks = torch.chunk(noise, noise.shape[0], dim=0)
# timesteps_chunks = torch.chunk(timesteps, timesteps.shape[0], dim=0)
# # Get the low res latents only if needed
# low_res_latents = self._get_low_res_latents(original_samples)
# low_res_latents_chunks = torch.chunk(low_res_latents, low_res_latents.shape[0], dim=0)
# # Handle case where timesteps is a single value for all samples
# if len(timesteps_chunks) == 1 and len(timesteps_chunks) != len(original_samples_chunks):
# timesteps_chunks = [timesteps_chunks[0]] * len(original_samples_chunks)
# noisy_latents_chunks = []
# effective_noise_chunks = [] # Store the effective noise for each sample
# for idx in range(original_samples.shape[0]):
# t = timesteps_chunks[idx]
# t_01 = (t / 1000).to(original_samples_chunks[idx].device)
# lrln = low_res_latents_chunks[idx] - original_samples_chunks[idx]
# # lrln = lrln * (1 - t_01)
# # make the noise an interpolation between noise and low_res_latents with
# # being noise at t_01=1 and low_res_latents at t_01=0
# new_noise = t_01 * noise_chunks[idx] + (1 - t_01) * lrln
# # new_noise = noise_chunks[idx] + lrln
# # new_noise = noise_chunks[idx] + lrln
# # Then apply flowmatching interpolation between this blended state and noise
# noisy_latents = (1 - t_01) * original_samples + t_01 * new_noise
# # For prediction target, we need to store the effective "source"
# effective_noise_chunks.append(new_noise)
# noisy_latents_chunks.append(noisy_latents)
# noisy_latents = torch.cat(noisy_latents_chunks, dim=0)
# self.effective_noise = torch.cat(effective_noise_chunks, dim=0) # Store for loss calculation
# return noisy_latents

View File

@@ -0,0 +1,272 @@
import inspect
import weakref
import torch
from typing import TYPE_CHECKING
from toolkit.lora_special import LoRASpecialNetwork
from diffusers import FluxTransformer2DModel
# weakref
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion
from toolkit.config_modules import AdapterConfig, TrainConfig, ModelConfig
from toolkit.custom_adapter import CustomAdapter
# after each step we concat the control image with the latents
# latent_model_input = torch.cat([latents, control_image], dim=2)
# the x_embedder has a full rank lora to handle the additional channels
# this replaces the x_embedder with a full rank lora. on flux this is
# x_embedder(diffusers) or img_in(bfl)
# Flux
# img_in.lora_A.weight [128, 128]
# img_in.lora_B.bias [3 072]
# img_in.lora_B.weight [3 072, 128]
class ImgEmbedder(torch.nn.Module):
def __init__(
self,
adapter: 'ControlLoraAdapter',
orig_layer: torch.nn.Linear,
in_channels=64,
out_channels=3072
):
super().__init__()
# only do the weight for the new input. We combine with the original linear layer
init = torch.randn(out_channels, in_channels, device=orig_layer.weight.device, dtype=orig_layer.weight.dtype) * 0.01
self.weight = torch.nn.Parameter(init)
self.adapter_ref: weakref.ref = weakref.ref(adapter)
self.orig_layer_ref: weakref.ref = weakref.ref(orig_layer)
@classmethod
def from_model(
cls,
model: FluxTransformer2DModel,
adapter: 'ControlLoraAdapter',
num_control_images=1,
has_inpainting_input=False
):
if model.__class__.__name__ == 'FluxTransformer2DModel':
num_adapter_in_channels = model.x_embedder.in_features * num_control_images
if has_inpainting_input:
# inpainting has the mask before packing latents. it is normally 16 ch + 1ch mask
# packed it is 64ch + 4ch mask
# so we need to add 4 to the input channels
num_adapter_in_channels += 4
x_embedder: torch.nn.Linear = model.x_embedder
img_embedder = cls(
adapter,
orig_layer=x_embedder,
in_channels=num_adapter_in_channels,
out_channels=x_embedder.out_features,
)
# hijack the forward method
x_embedder._orig_ctrl_lora_forward = x_embedder.forward
x_embedder.forward = img_embedder.forward
# update the config of the transformer
model.config.in_channels = model.config.in_channels * (num_control_images + 1)
model.config["in_channels"] = model.config.in_channels
return img_embedder
else:
raise ValueError("Model not supported")
@property
def is_active(self):
return self.adapter_ref().is_active
def forward(self, x):
if not self.is_active:
# make sure lora is not active
if self.adapter_ref().control_lora is not None:
self.adapter_ref().control_lora.is_active = False
return self.orig_layer_ref()._orig_ctrl_lora_forward(x)
# make sure lora is active
if self.adapter_ref().control_lora is not None:
self.adapter_ref().control_lora.is_active = True
orig_device = x.device
orig_dtype = x.dtype
x = x.to(self.weight.device, dtype=self.weight.dtype)
orig_weight = self.orig_layer_ref().weight.data.detach()
orig_weight = orig_weight.to(self.weight.device, dtype=self.weight.dtype)
linear_weight = torch.cat([orig_weight, self.weight], dim=1)
bias = None
if self.orig_layer_ref().bias is not None:
bias = self.orig_layer_ref().bias.data.detach().to(self.weight.device, dtype=self.weight.dtype)
x = torch.nn.functional.linear(x, linear_weight, bias)
x = x.to(orig_device, dtype=orig_dtype)
return x
class ControlLoraAdapter(torch.nn.Module):
def __init__(
self,
adapter: 'CustomAdapter',
sd: 'StableDiffusion',
config: 'AdapterConfig',
train_config: 'TrainConfig'
):
super().__init__()
self.adapter_ref: weakref.ref = weakref.ref(adapter)
self.sd_ref = weakref.ref(sd)
self.model_config: ModelConfig = sd.model_config
self.network_config = config.lora_config
self.train_config = train_config
self.device_torch = sd.device_torch
self.control_lora = None
if self.network_config is not None:
network_kwargs = {} if self.network_config.network_kwargs is None else self.network_config.network_kwargs
if hasattr(sd, 'target_lora_modules'):
network_kwargs['target_lin_modules'] = self.sd.target_lora_modules
if 'ignore_if_contains' not in network_kwargs:
network_kwargs['ignore_if_contains'] = []
# always ignore x_embedder
network_kwargs['ignore_if_contains'].append('x_embedder')
self.control_lora = LoRASpecialNetwork(
text_encoder=sd.text_encoder,
unet=sd.unet,
lora_dim=self.network_config.linear,
multiplier=1.0,
alpha=self.network_config.linear_alpha,
train_unet=self.train_config.train_unet,
train_text_encoder=self.train_config.train_text_encoder,
conv_lora_dim=self.network_config.conv,
conv_alpha=self.network_config.conv_alpha,
is_sdxl=self.model_config.is_xl or self.model_config.is_ssd,
is_v2=self.model_config.is_v2,
is_v3=self.model_config.is_v3,
is_pixart=self.model_config.is_pixart,
is_auraflow=self.model_config.is_auraflow,
is_flux=self.model_config.is_flux,
is_lumina2=self.model_config.is_lumina2,
is_ssd=self.model_config.is_ssd,
is_vega=self.model_config.is_vega,
dropout=self.network_config.dropout,
use_text_encoder_1=self.model_config.use_text_encoder_1,
use_text_encoder_2=self.model_config.use_text_encoder_2,
use_bias=False,
is_lorm=False,
network_config=self.network_config,
network_type=self.network_config.type,
transformer_only=self.network_config.transformer_only,
is_transformer=sd.is_transformer,
base_model=sd,
**network_kwargs
)
self.control_lora.force_to(self.device_torch, dtype=torch.float32)
self.control_lora._update_torch_multiplier()
self.control_lora.apply_to(
sd.text_encoder,
sd.unet,
self.train_config.train_text_encoder,
self.train_config.train_unet
)
self.control_lora.can_merge_in = False
self.control_lora.prepare_grad_etc(sd.text_encoder, sd.unet)
if self.train_config.gradient_checkpointing:
self.control_lora.enable_gradient_checkpointing()
self.x_embedder = ImgEmbedder.from_model(
sd.unet,
self,
num_control_images=config.num_control_images,
has_inpainting_input=config.has_inpainting_input
)
self.x_embedder.to(self.device_torch)
def get_params(self):
if self.control_lora is not None:
config = {
'text_encoder_lr': self.train_config.lr,
'unet_lr': self.train_config.lr,
}
sig = inspect.signature(self.control_lora.prepare_optimizer_params)
if 'default_lr' in sig.parameters:
config['default_lr'] = self.train_config.lr
if 'learning_rate' in sig.parameters:
config['learning_rate'] = self.train_config.lr
params_net = self.control_lora.prepare_optimizer_params(
**config
)
# we want only tensors here
params = []
for p in params_net:
if isinstance(p, dict):
params += p["params"]
elif isinstance(p, torch.Tensor):
params.append(p)
elif isinstance(p, list):
params += p
else:
params = []
# make sure the embedder is float32
self.x_embedder.to(torch.float32)
params += list(self.x_embedder.parameters())
# we need to be able to yield from the list like yield from params
return params
def load_weights(self, state_dict, strict=True):
lora_sd = {}
img_embedder_sd = {}
for key, value in state_dict.items():
if "x_embedder" in key:
new_key = key.replace("transformer.x_embedder.", "")
img_embedder_sd[new_key] = value
else:
lora_sd[key] = value
# todo process state dict before loading
if self.control_lora is not None:
self.control_lora.load_weights(lora_sd)
# automatically upgrade the x imbedder if more dims are added
if self.x_embedder.weight.shape[1] > img_embedder_sd['weight'].shape[1]:
print("Upgrading x_embedder from {} to {}".format(
img_embedder_sd['weight'].shape[1],
self.x_embedder.weight.shape[1]
))
while img_embedder_sd['weight'].shape[1] < self.x_embedder.weight.shape[1]:
img_embedder_sd['weight'] = torch.cat([img_embedder_sd['weight'] ] * 2, dim=1)
if img_embedder_sd['weight'].shape[1] > self.x_embedder.weight.shape[1]:
img_embedder_sd['weight'] = img_embedder_sd['weight'][:, :self.x_embedder.weight.shape[1]]
self.x_embedder.load_state_dict(img_embedder_sd, strict=False)
def get_state_dict(self):
if self.control_lora is not None:
lora_sd = self.control_lora.get_state_dict(dtype=torch.float32)
else:
lora_sd = {}
# todo make sure we match loras elseware.
img_embedder_sd = self.x_embedder.state_dict()
for key, value in img_embedder_sd.items():
lora_sd[f"transformer.x_embedder.{key}"] = value
return lora_sd
@property
def is_active(self):
return self.adapter_ref().is_active

View File

@@ -0,0 +1,33 @@
import torch
class Decorator(torch.nn.Module):
def __init__(
self,
num_tokens: int = 4,
token_size: int = 4096,
) -> None:
super().__init__()
self.weight: torch.nn.Parameter = torch.nn.Parameter(
torch.randn(num_tokens, token_size)
)
# ensure it is float32
self.weight.data = self.weight.data.float()
def forward(self, text_embeds: torch.Tensor, is_unconditional=False) -> torch.Tensor:
# make sure the param is float32
if self.weight.dtype != text_embeds.dtype:
self.weight.data = self.weight.data.float()
# expand batch to match text_embeds
batch_size = text_embeds.shape[0]
decorator_embeds = self.weight.unsqueeze(0).expand(batch_size, -1, -1)
if is_unconditional:
# zero pad the decorator embeds
decorator_embeds = torch.zeros_like(decorator_embeds)
if decorator_embeds.dtype != text_embeds.dtype:
decorator_embeds = decorator_embeds.to(text_embeds.dtype)
text_embeds = torch.cat((text_embeds, decorator_embeds), dim=-2)
return text_embeds

View File

@@ -0,0 +1,367 @@
import torch
import os
from torch import nn
from safetensors.torch import load_file
import torch.nn.functional as F
from diffusers import AutoencoderTiny
from transformers import SiglipImageProcessor, SiglipVisionModel
import lpips
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
class ResBlock(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
self.norm1 = nn.GroupNorm(8, out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
self.norm2 = nn.GroupNorm(8, out_channels)
self.skip = nn.Conv2d(in_channels, out_channels,
1) if in_channels != out_channels else nn.Identity()
def forward(self, x):
identity = self.skip(x)
x = self.conv1(x)
x = self.norm1(x)
x = F.silu(x)
x = self.conv2(x)
x = self.norm2(x)
x = F.silu(x + identity)
return x
class DiffusionFeatureExtractor2(nn.Module):
def __init__(self, in_channels=32):
super().__init__()
self.version = 2
# Path 1: Upsample to 512x512 (1, 64, 512, 512)
self.up_path = nn.ModuleList([
nn.Conv2d(in_channels, 64, 3, padding=1),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(64, 64),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(64, 64),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(64, 64),
nn.Conv2d(64, 64, 3, padding=1),
])
# Path 2: Upsample to 256x256 (1, 128, 256, 256)
self.path2 = nn.ModuleList([
nn.Conv2d(in_channels, 128, 3, padding=1),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(128, 128),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(128, 128),
nn.Conv2d(128, 128, 3, padding=1),
])
# Path 3: Upsample to 128x128 (1, 256, 128, 128)
self.path3 = nn.ModuleList([
nn.Conv2d(in_channels, 256, 3, padding=1),
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
ResBlock(256, 256),
nn.Conv2d(256, 256, 3, padding=1)
])
# Path 4: Original size (1, 512, 64, 64)
self.path4 = nn.ModuleList([
nn.Conv2d(in_channels, 512, 3, padding=1),
ResBlock(512, 512),
ResBlock(512, 512),
nn.Conv2d(512, 512, 3, padding=1)
])
# Path 5: Downsample to 32x32 (1, 512, 32, 32)
self.path5 = nn.ModuleList([
nn.Conv2d(in_channels, 512, 3, padding=1),
ResBlock(512, 512),
nn.AvgPool2d(2),
ResBlock(512, 512),
nn.Conv2d(512, 512, 3, padding=1)
])
def forward(self, x):
outputs = []
# Path 1: 512x512
x1 = x
for layer in self.up_path:
x1 = layer(x1)
outputs.append(x1) # [1, 64, 512, 512]
# Path 2: 256x256
x2 = x
for layer in self.path2:
x2 = layer(x2)
outputs.append(x2) # [1, 128, 256, 256]
# Path 3: 128x128
x3 = x
for layer in self.path3:
x3 = layer(x3)
outputs.append(x3) # [1, 256, 128, 128]
# Path 4: 64x64
x4 = x
for layer in self.path4:
x4 = layer(x4)
outputs.append(x4) # [1, 512, 64, 64]
# Path 5: 32x32
x5 = x
for layer in self.path5:
x5 = layer(x5)
outputs.append(x5) # [1, 512, 32, 32]
return outputs
class DFEBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)
self.act = nn.GELU()
def forward(self, x):
x_in = x
x = self.conv1(x)
x = self.conv2(x)
x = self.act(x)
x = x + x_in
return x
class DiffusionFeatureExtractor(nn.Module):
def __init__(self, in_channels=32):
super().__init__()
self.version = 1
num_blocks = 6
self.conv_in = nn.Conv2d(in_channels, 512, 1)
self.blocks = nn.ModuleList([DFEBlock(512) for _ in range(num_blocks)])
self.conv_out = nn.Conv2d(512, 512, 1)
def forward(self, x):
x = self.conv_in(x)
for block in self.blocks:
x = block(x)
x = self.conv_out(x)
return x
class DiffusionFeatureExtractor3(nn.Module):
def __init__(self, device=torch.device("cuda"), dtype=torch.bfloat16):
super().__init__()
self.version = 3
vae = AutoencoderTiny.from_pretrained(
"madebyollin/taef1", torch_dtype=torch.bfloat16)
self.vae = vae
image_encoder_path = "google/siglip-so400m-patch14-384"
try:
self.image_processor = SiglipImageProcessor.from_pretrained(
image_encoder_path)
except EnvironmentError:
self.image_processor = SiglipImageProcessor()
self.vision_encoder = SiglipVisionModel.from_pretrained(
image_encoder_path,
ignore_mismatched_sizes=True
).to(device, dtype=dtype)
self.lpips_model = lpips_model = lpips.LPIPS(net='vgg')
self.lpips_model = lpips_model.to(device, dtype=torch.float32)
self.losses = {}
self.log_every = 100
self.step = 0
def get_siglip_features(self, tensors_0_1):
dtype = torch.bfloat16
device = self.vae.device
# resize to 384x384
images = F.interpolate(tensors_0_1, size=(384, 384),
mode='bicubic', align_corners=False)
mean = torch.tensor(self.image_processor.image_mean).to(
device, dtype=dtype
).detach()
std = torch.tensor(self.image_processor.image_std).to(
device, dtype=dtype
).detach()
# tensors_0_1 = torch.clip((255. * tensors_0_1), 0, 255).round() / 255.0
clip_image = (
images - mean.view([1, 3, 1, 1])) / std.view([1, 3, 1, 1])
id_embeds = self.vision_encoder(
clip_image,
output_hidden_states=True,
)
last_hidden_state = id_embeds['last_hidden_state']
return last_hidden_state
def get_lpips_features(self, tensors_0_1):
device = self.vae.device
tensors_n1p1 = (tensors_0_1 * 2) - 1
def get_lpips_features(img): # -1 to 1
in0_input = self.lpips_model.scaling_layer(img)
outs0 = self.lpips_model.net.forward(in0_input)
feats0 = {}
feats_list = []
for kk in range(self.lpips_model.L):
feats0[kk] = lpips.normalize_tensor(outs0[kk])
feats_list.append(feats0[kk])
# 512 in
# vgg
# 0 torch.Size([1, 64, 512, 512])
# 1 torch.Size([1, 128, 256, 256])
# 2 torch.Size([1, 256, 128, 128])
# 3 torch.Size([1, 512, 64, 64])
# 4 torch.Size([1, 512, 32, 32])
return feats_list
# do lpips
lpips_feat_list = [x for x in get_lpips_features(
tensors_n1p1.to(device, dtype=torch.float32))]
return lpips_feat_list
def forward(
self,
noise,
noise_pred,
noisy_latents,
timesteps,
batch: DataLoaderBatchDTO,
scheduler: CustomFlowMatchEulerDiscreteScheduler,
# lpips_weight=1.0,
lpips_weight=10.0,
clip_weight=0.1,
pixel_weight=0.1
):
dtype = torch.bfloat16
device = self.vae.device
# first we step the scheduler from current timestep to the very end for a full denoise
# bs = noise_pred.shape[0]
# noise_pred_chunks = torch.chunk(noise_pred, bs)
# timestep_chunks = torch.chunk(timesteps, bs)
# noisy_latent_chunks = torch.chunk(noisy_latents, bs)
# stepped_chunks = []
# for idx in range(bs):
# model_output = noise_pred_chunks[idx]
# timestep = timestep_chunks[idx]
# scheduler._step_index = None
# scheduler._init_step_index(timestep)
# sample = noisy_latent_chunks[idx].to(torch.float32)
# sigma = scheduler.sigmas[scheduler.step_index]
# sigma_next = scheduler.sigmas[-1] # use last sigma for final step
# prev_sample = sample + (sigma_next - sigma) * model_output
# stepped_chunks.append(prev_sample)
# stepped_latents = torch.cat(stepped_chunks, dim=0)
stepped_latents = noise - noise_pred
latents = stepped_latents.to(self.vae.device, dtype=self.vae.dtype)
latents = (
latents / self.vae.config['scaling_factor']) + self.vae.config['shift_factor']
tensors_n1p1 = self.vae.decode(latents).sample # -1 to 1
pred_images = (tensors_n1p1 + 1) / 2 # 0 to 1
lpips_feat_list_pred = self.get_lpips_features(pred_images.float())
total_loss = 0
with torch.no_grad():
target_img = batch.tensor.to(device, dtype=dtype)
# go from -1 to 1 to 0 to 1
target_img = (target_img + 1) / 2
lpips_feat_list_target = self.get_lpips_features(target_img.float())
if clip_weight > 0:
target_clip_output = self.get_siglip_features(target_img).detach()
if clip_weight > 0:
pred_clip_output = self.get_siglip_features(pred_images)
clip_loss = torch.nn.functional.mse_loss(
pred_clip_output.float(), target_clip_output.float()
) * clip_weight
if 'clip_loss' not in self.losses:
self.losses['clip_loss'] = clip_loss.item()
else:
self.losses['clip_loss'] += clip_loss.item()
total_loss += clip_loss
skip_lpips_layers = []
lpips_loss = 0
for idx, lpips_feat in enumerate(lpips_feat_list_pred):
if idx in skip_lpips_layers:
continue
lpips_loss += torch.nn.functional.mse_loss(
lpips_feat.float(), lpips_feat_list_target[idx].float()
) * lpips_weight
if f'lpips_loss_{idx}' not in self.losses:
self.losses[f'lpips_loss_{idx}'] = lpips_loss.item()
else:
self.losses[f'lpips_loss_{idx}'] += lpips_loss.item()
total_loss += lpips_loss
# mse_loss = torch.nn.functional.mse_loss(
# stepped_latents.float(), batch.latents.float()
# ) * pixel_weight
# if 'pixel_loss' not in self.losses:
# self.losses['pixel_loss'] = mse_loss.item()
# else:
# self.losses['pixel_loss'] += mse_loss.item()
if self.step % self.log_every == 0 and self.step > 0:
print(f"DFE losses:")
for key in self.losses:
self.losses[key] /= self.log_every
# print in 2.000e-01 format
print(f" - {key}: {self.losses[key]:.3e}")
self.losses[key] = 0.0
# total_loss += mse_loss
self.step += 1
return total_loss
def load_dfe(model_path) -> DiffusionFeatureExtractor:
if model_path == "v3":
dfe = DiffusionFeatureExtractor3()
dfe.eval()
return dfe
if not os.path.exists(model_path):
raise FileNotFoundError(f"Model file not found: {model_path}")
# if it ende with safetensors
if model_path.endswith('.safetensors'):
state_dict = load_file(model_path)
else:
state_dict = torch.load(model_path, weights_only=True)
if 'model_state_dict' in state_dict:
state_dict = state_dict['model_state_dict']
if 'conv_in.weight' in state_dict:
dfe = DiffusionFeatureExtractor()
else:
dfe = DiffusionFeatureExtractor2()
dfe.load_state_dict(state_dict)
dfe.eval()
return dfe

993
toolkit/models/flex2.py Normal file
View File

@@ -0,0 +1,993 @@
from typing import List, Optional, Union
from diffusers import FluxPipeline
import inspect
from typing import Any, Callable, Dict, List, Optional, Union
import numpy as np
import torch
from diffusers.loaders import FluxLoraLoaderMixin, TextualInversionLoaderMixin
from diffusers.utils import (
USE_PEFT_BACKEND,
is_torch_xla_available,
logging,
replace_example_docstring,
scale_lora_layers,
unscale_lora_layers,
)
from transformers import AutoModel, AutoTokenizer
from transformers import (
CLIPImageProcessor,
CLIPTextModel,
CLIPTokenizer,
CLIPVisionModelWithProjection
)
from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
from diffusers.loaders import FluxIPAdapterMixin, FluxLoraLoaderMixin, FromSingleFileMixin, TextualInversionLoaderMixin
from diffusers.models import AutoencoderKL, FluxTransformer2DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils.torch_utils import randn_tensor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.flux.pipeline_output import FluxPipelineOutput
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
EXAMPLE_DOC_STRING = """
Examples:
```py
>>> import torch
>>> from diffusers import Flex2Pipeline
>>> pipe = Flex2Pipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16)
>>> pipe.to("cuda")
>>> prompt = "A cat holding a sign that says hello world"
>>> # Depending on the variant being used, the pipeline call will slightly vary.
>>> # Refer to the pipeline documentation for more details.
>>> image = pipe(prompt, num_inference_steps=4, guidance_scale=0.0).images[0]
>>> image.save("flux.png")
```
"""
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def calculate_shift(
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.16,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
def retrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] = None,
device: Optional[Union[str, torch.device]] = None,
timesteps: Optional[List[int]] = None,
sigmas: Optional[List[float]] = None,
**kwargs,
):
r"""
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
Args:
scheduler (`SchedulerMixin`):
The scheduler to get timesteps from.
num_inference_steps (`int`):
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
must be `None`.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
timesteps (`List[int]`, *optional*):
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
`num_inference_steps` and `sigmas` must be `None`.
sigmas (`List[float]`, *optional*):
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
`num_inference_steps` and `timesteps` must be `None`.
Returns:
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
second element is the number of inference steps.
"""
if timesteps is not None and sigmas is not None:
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
if timesteps is not None:
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accepts_timesteps:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
elif sigmas is not None:
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
if not accept_sigmas:
raise ValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps = scheduler.timesteps
num_inference_steps = len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps = scheduler.timesteps
return timesteps, num_inference_steps
class Flex2Pipeline(
DiffusionPipeline,
FluxLoraLoaderMixin,
FromSingleFileMixin,
TextualInversionLoaderMixin,
FluxIPAdapterMixin,
):
r"""
The Flux pipeline for text-to-image generation.
Reference: https://blackforestlabs.ai/announcing-black-forest-labs/
Args:
transformer ([`FluxTransformer2DModel`]):
Conditional Transformer (MMDiT) architecture to denoise the encoded image latents.
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
vae ([`AutoencoderKL`]):
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
text_encoder ([`CLIPTextModel`]):
[CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
text_encoder_2 ([`T5EncoderModel`]):
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
tokenizer (`CLIPTokenizer`):
Tokenizer of class
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
tokenizer_2 (`T5TokenizerFast`):
Second Tokenizer of class
[T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
"""
model_cpu_offload_seq = "text_encoder->text_encoder_2->image_encoder->transformer->vae"
_optional_components = ["image_encoder", "feature_extractor"]
_callback_tensor_inputs = ["latents", "prompt_embeds"]
def __init__(
self,
scheduler: FlowMatchEulerDiscreteScheduler,
vae: AutoencoderKL,
text_encoder: CLIPTextModel,
tokenizer: CLIPTokenizer,
text_encoder_2: AutoModel,
tokenizer_2: AutoTokenizer,
transformer: FluxTransformer2DModel,
image_encoder: CLIPVisionModelWithProjection = None,
feature_extractor: CLIPImageProcessor = None,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
tokenizer=tokenizer,
tokenizer_2=tokenizer_2,
transformer=transformer,
scheduler=scheduler,
image_encoder=image_encoder,
feature_extractor=feature_extractor,
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) if getattr(self, "vae", None) else 8
# Flux latents are turned into 2x2 patches and packed. This means the latent width and height has to be divisible
# by the patch size. So the vae scale factor is multiplied by the patch size to account for this
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * 2)
self.tokenizer_max_length = (
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
)
self.default_sample_size = 128
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
# determine length of system prompt
self.system_prompt_length = self.tokenizer_2(
[self.system_prompt],
padding="longest",
return_tensors="pt",
).input_ids[0].shape[0]
def _get_clip_prompt_embeds(
self,
prompt: Union[str, List[str]],
num_images_per_prompt: int = 1,
device: Optional[torch.device] = None,
):
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
if isinstance(self, TextualInversionLoaderMixin):
prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.tokenizer_max_length,
truncation=True,
return_overflowing_tokens=False,
return_length=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because CLIP can only handle sequences up to"
f" {self.tokenizer_max_length} tokens: {removed_text}"
)
prompt_embeds = self.text_encoder(text_input_ids.to(device), output_hidden_states=False)
# Use pooled output of CLIPTextModel
prompt_embeds = prompt_embeds.pooler_output
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
return prompt_embeds
def _get_llm_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_images_per_prompt: int = 1,
max_sequence_length: int = 512,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
if isinstance(self, TextualInversionLoaderMixin):
prompt = self.maybe_convert_prompt(prompt, self.tokenizer_2)
text_inputs = self.tokenizer_2(
prompt,
padding="max_length",
max_length=max_sequence_length + self.system_prompt_length,
truncation=True,
return_length=False,
return_overflowing_tokens=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids.to(device)
prompt_attention_mask = text_inputs.attention_mask.to(device)
untruncated_ids = self.tokenizer_2(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer_2.batch_decode(untruncated_ids[:, self.tokenizer_max_length - 1 : -1])
logger.warning(
"The following part of your input was truncated because `max_sequence_length` is set to "
f" {max_sequence_length + self.system_prompt_length} tokens: {removed_text}"
)
prompt_embeds = self.text_encoder_2(
text_input_ids,
attention_mask=prompt_attention_mask,
output_hidden_states=True
)
prompt_embeds = prompt_embeds.hidden_states[-1]
# remove the system prompt from the input and attention mask
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
dtype = self.text_encoder_2.dtype
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
_, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
return prompt_embeds
def encode_prompt(
self,
prompt: Union[str, List[str]],
prompt_2: Union[str, List[str]],
device: Optional[torch.device] = None,
num_images_per_prompt: int = 1,
prompt_embeds: Optional[torch.FloatTensor] = None,
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
max_sequence_length: int = 512,
lora_scale: Optional[float] = None,
):
r"""
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
prompt_2 (`str` or `List[str]`, *optional*):
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
used in all text-encoders
device: (`torch.device`):
torch device
num_images_per_prompt (`int`):
number of images that should be generated per prompt
prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
If not provided, pooled text embeddings will be generated from `prompt` input argument.
lora_scale (`float`, *optional*):
A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded.
"""
device = device or self._execution_device
# set lora scale so that monkey patched LoRA
# function of text encoder can correctly access it
if lora_scale is not None and isinstance(self, FluxLoraLoaderMixin):
self._lora_scale = lora_scale
# dynamically adjust the LoRA scale
if self.text_encoder is not None and USE_PEFT_BACKEND:
scale_lora_layers(self.text_encoder, lora_scale)
if self.text_encoder_2 is not None and USE_PEFT_BACKEND:
scale_lora_layers(self.text_encoder_2, lora_scale)
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt_embeds is None:
prompt_2 = prompt_2 or prompt
prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2
# We only use the pooled prompt output from the CLIPTextModel
pooled_prompt_embeds = self._get_clip_prompt_embeds(
prompt=prompt,
device=device,
num_images_per_prompt=num_images_per_prompt,
)
prompt_embeds = self._get_llm_prompt_embeds(
prompt=prompt_2,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
)
if self.text_encoder is not None:
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
# Retrieve the original scale by scaling back the LoRA layers
unscale_lora_layers(self.text_encoder, lora_scale)
if self.text_encoder_2 is not None:
if isinstance(self, FluxLoraLoaderMixin) and USE_PEFT_BACKEND:
# Retrieve the original scale by scaling back the LoRA layers
unscale_lora_layers(self.text_encoder_2, lora_scale)
dtype = self.text_encoder.dtype if self.text_encoder is not None else self.transformer.dtype
text_ids = torch.zeros(prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
return prompt_embeds, pooled_prompt_embeds, text_ids
def encode_image(self, image, device, num_images_per_prompt):
dtype = next(self.image_encoder.parameters()).dtype
if not isinstance(image, torch.Tensor):
image = self.feature_extractor(image, return_tensors="pt").pixel_values
image = image.to(device=device, dtype=dtype)
image_embeds = self.image_encoder(image).image_embeds
image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0)
return image_embeds
def prepare_ip_adapter_image_embeds(
self, ip_adapter_image, ip_adapter_image_embeds, device, num_images_per_prompt
):
image_embeds = []
if ip_adapter_image_embeds is None:
if not isinstance(ip_adapter_image, list):
ip_adapter_image = [ip_adapter_image]
if len(ip_adapter_image) != len(self.transformer.encoder_hid_proj.image_projection_layers):
raise ValueError(
f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(self.transformer.encoder_hid_proj.image_projection_layers)} IP Adapters."
)
for single_ip_adapter_image, image_proj_layer in zip(
ip_adapter_image, self.transformer.encoder_hid_proj.image_projection_layers
):
single_image_embeds = self.encode_image(single_ip_adapter_image, device, 1)
image_embeds.append(single_image_embeds[None, :])
else:
for single_image_embeds in ip_adapter_image_embeds:
image_embeds.append(single_image_embeds)
ip_adapter_image_embeds = []
for i, single_image_embeds in enumerate(image_embeds):
single_image_embeds = torch.cat([single_image_embeds] * num_images_per_prompt, dim=0)
single_image_embeds = single_image_embeds.to(device=device)
ip_adapter_image_embeds.append(single_image_embeds)
return ip_adapter_image_embeds
def check_inputs(
self,
prompt,
prompt_2,
height,
width,
negative_prompt=None,
negative_prompt_2=None,
prompt_embeds=None,
negative_prompt_embeds=None,
pooled_prompt_embeds=None,
negative_pooled_prompt_embeds=None,
callback_on_step_end_tensor_inputs=None,
max_sequence_length=None,
):
if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0:
logger.warning(
f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly"
)
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt_2 is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)):
raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
elif negative_prompt_2 is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt_2`: {negative_prompt_2} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
if prompt_embeds is not None and pooled_prompt_embeds is None:
raise ValueError(
"If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`."
)
if negative_prompt_embeds is not None and negative_pooled_prompt_embeds is None:
raise ValueError(
"If `negative_prompt_embeds` are provided, `negative_pooled_prompt_embeds` also have to be passed. Make sure to generate `negative_pooled_prompt_embeds` from the same text encoder that was used to generate `negative_prompt_embeds`."
)
if max_sequence_length is not None and max_sequence_length > 512:
raise ValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}")
@staticmethod
def _prepare_latent_image_ids(batch_size, height, width, device, dtype):
latent_image_ids = torch.zeros(height, width, 3)
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape
latent_image_ids = latent_image_ids.reshape(
latent_image_id_height * latent_image_id_width, latent_image_id_channels
)
return latent_image_ids.to(device=device, dtype=dtype)
@staticmethod
def _pack_latents(latents, batch_size, num_channels_latents, height, width):
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
latents = latents.permute(0, 2, 4, 1, 3, 5)
latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
return latents
@staticmethod
def _unpack_latents(latents, height, width, vae_scale_factor):
batch_size, num_patches, channels = latents.shape
# VAE applies 8x compression on images but we must also account for packing which requires
# latent height and width to be divisible by 2.
height = 2 * (int(height) // (vae_scale_factor * 2))
width = 2 * (int(width) // (vae_scale_factor * 2))
latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2)
latents = latents.permute(0, 3, 1, 4, 2, 5)
latents = latents.reshape(batch_size, channels // (2 * 2), height, width)
return latents
def enable_vae_slicing(self):
r"""
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
"""
self.vae.enable_slicing()
def disable_vae_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
computing decoding in one step.
"""
self.vae.disable_slicing()
def enable_vae_tiling(self):
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
processing larger images.
"""
self.vae.enable_tiling()
def disable_vae_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
computing decoding in one step.
"""
self.vae.disable_tiling()
def prepare_latents(
self,
batch_size,
num_channels_latents,
height,
width,
dtype,
device,
generator,
latents=None,
):
# VAE applies 8x compression on images but we must also account for packing which requires
# latent height and width to be divisible by 2.
height = 2 * (int(height) // (self.vae_scale_factor * 2))
width = 2 * (int(width) // (self.vae_scale_factor * 2))
shape = (batch_size, num_channels_latents, height, width)
if latents is not None:
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
return latents.to(device=device, dtype=dtype), latent_image_ids
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width)
latent_image_ids = self._prepare_latent_image_ids(batch_size, height // 2, width // 2, device, dtype)
return latents, latent_image_ids
@property
def guidance_scale(self):
return self._guidance_scale
@property
def joint_attention_kwargs(self):
return self._joint_attention_kwargs
@property
def num_timesteps(self):
return self._num_timesteps
@property
def current_timestep(self):
return self._current_timestep
@property
def interrupt(self):
return self._interrupt
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
prompt_2: Optional[Union[str, List[str]]] = None,
negative_prompt: Union[str, List[str]] = None,
negative_prompt_2: Optional[Union[str, List[str]]] = None,
true_cfg_scale: float = 1.0,
height: Optional[int] = None,
width: Optional[int] = None,
num_inference_steps: int = 28,
sigmas: Optional[List[float]] = None,
guidance_scale: float = 3.5,
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
ip_adapter_image: Optional[PipelineImageInput] = None,
ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
negative_ip_adapter_image: Optional[PipelineImageInput] = None,
negative_ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: Optional[str] = "pil",
return_dict: bool = True,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
):
r"""
Function invoked when calling the pipeline for generation.
Args:
prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
instead.
prompt_2 (`str` or `List[str]`, *optional*):
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
will be used instead.
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `true_cfg_scale` is
not greater than `1`).
negative_prompt_2 (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
`text_encoder_2`. If not defined, `negative_prompt` is used in all the text-encoders.
true_cfg_scale (`float`, *optional*, defaults to 1.0):
When > 1.0 and a provided `negative_prompt`, enables true classifier-free guidance.
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
The height in pixels of the generated image. This is set to 1024 by default for the best results.
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
The width in pixels of the generated image. This is set to 1024 by default for the best results.
num_inference_steps (`int`, *optional*, defaults to 50):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
sigmas (`List[float]`, *optional*):
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
will be used.
guidance_scale (`float`, *optional*, defaults to 7.0):
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
`guidance_scale` is defined as `w` of equation 2. of [Imagen
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
usually at the expense of lower image quality.
num_images_per_prompt (`int`, *optional*, defaults to 1):
The number of images to generate per prompt.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
to make generation deterministic.
latents (`torch.FloatTensor`, *optional*):
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor will ge generated by sampling using the supplied random `generator`.
prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
If not provided, pooled text embeddings will be generated from `prompt` input argument.
ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
provided, embeddings are computed from the `ip_adapter_image` input argument.
negative_ip_adapter_image:
(`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters.
negative_ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*):
Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of
IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not
provided, embeddings are computed from the `ip_adapter_image` input argument.
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
input argument.
output_type (`str`, *optional*, defaults to `"pil"`):
The output format of the generate image. Choose between
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple.
joint_attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
`self.processor` in
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
callback_on_step_end (`Callable`, *optional*):
A function that calls at the end of each denoising steps during the inference. The function is called
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
`callback_on_step_end_tensor_inputs`.
callback_on_step_end_tensor_inputs (`List`, *optional*):
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
`._callback_tensor_inputs` attribute of your pipeline class.
max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`.
Examples:
Returns:
[`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict`
is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated
images.
"""
height = height or self.default_sample_size * self.vae_scale_factor
width = width or self.default_sample_size * self.vae_scale_factor
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt,
prompt_2,
height,
width,
negative_prompt=negative_prompt,
negative_prompt_2=negative_prompt_2,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
pooled_prompt_embeds=pooled_prompt_embeds,
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
max_sequence_length=max_sequence_length,
)
self._guidance_scale = guidance_scale
self._joint_attention_kwargs = joint_attention_kwargs
self._current_timestep = None
self._interrupt = False
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
lora_scale = (
self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None
)
has_neg_prompt = negative_prompt is not None or (
negative_prompt_embeds is not None and negative_pooled_prompt_embeds is not None
)
do_true_cfg = true_cfg_scale > 1 and has_neg_prompt
(
prompt_embeds,
pooled_prompt_embeds,
text_ids,
) = self.encode_prompt(
prompt=prompt,
prompt_2=prompt_2,
prompt_embeds=prompt_embeds,
pooled_prompt_embeds=pooled_prompt_embeds,
device=device,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
lora_scale=lora_scale,
)
if do_true_cfg:
(
negative_prompt_embeds,
negative_pooled_prompt_embeds,
_,
) = self.encode_prompt(
prompt=negative_prompt,
prompt_2=negative_prompt_2,
prompt_embeds=negative_prompt_embeds,
pooled_prompt_embeds=negative_pooled_prompt_embeds,
device=device,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
lora_scale=lora_scale,
)
# 4. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels // 4
latents, latent_image_ids = self.prepare_latents(
batch_size * num_images_per_prompt,
num_channels_latents,
height,
width,
prompt_embeds.dtype,
device,
generator,
latents,
)
# 5. Prepare timesteps
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas
image_seq_len = latents.shape[1]
mu = calculate_shift(
image_seq_len,
self.scheduler.config.get("base_image_seq_len", 256),
self.scheduler.config.get("max_image_seq_len", 4096),
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.16),
)
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
device,
sigmas=sigmas,
mu=mu,
)
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
self._num_timesteps = len(timesteps)
# handle guidance
if self.transformer.config.guidance_embeds:
guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32)
guidance = guidance.expand(latents.shape[0])
else:
guidance = None
if (ip_adapter_image is not None or ip_adapter_image_embeds is not None) and (
negative_ip_adapter_image is None and negative_ip_adapter_image_embeds is None
):
negative_ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
elif (ip_adapter_image is None and ip_adapter_image_embeds is None) and (
negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None
):
ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8)
if self.joint_attention_kwargs is None:
self._joint_attention_kwargs = {}
image_embeds = None
negative_image_embeds = None
if ip_adapter_image is not None or ip_adapter_image_embeds is not None:
image_embeds = self.prepare_ip_adapter_image_embeds(
ip_adapter_image,
ip_adapter_image_embeds,
device,
batch_size * num_images_per_prompt,
)
if negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None:
negative_image_embeds = self.prepare_ip_adapter_image_embeds(
negative_ip_adapter_image,
negative_ip_adapter_image_embeds,
device,
batch_size * num_images_per_prompt,
)
# 6. Denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
self._current_timestep = t
if image_embeds is not None:
self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latents.shape[0]).to(latents.dtype)
noise_pred = self.transformer(
hidden_states=latents,
timestep=timestep / 1000,
guidance=guidance,
pooled_projections=pooled_prompt_embeds,
encoder_hidden_states=prompt_embeds,
txt_ids=text_ids,
img_ids=latent_image_ids,
joint_attention_kwargs=self.joint_attention_kwargs,
return_dict=False,
)[0]
if do_true_cfg:
if negative_image_embeds is not None:
self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds
neg_noise_pred = self.transformer(
hidden_states=latents,
timestep=timestep / 1000,
guidance=guidance,
pooled_projections=negative_pooled_prompt_embeds,
encoder_hidden_states=negative_prompt_embeds,
txt_ids=text_ids,
img_ids=latent_image_ids,
joint_attention_kwargs=self.joint_attention_kwargs,
return_dict=False,
)[0]
noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred)
# compute the previous noisy sample x_t -> x_t-1
latents_dtype = latents.dtype
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
if latents.dtype != latents_dtype:
if torch.backends.mps.is_available():
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
latents = latents.to(latents_dtype)
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
self._current_timestep = None
if output_type == "latent":
image = latents
else:
latents = self._unpack_latents(latents, height, width, self.vae_scale_factor)
latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor
image = self.vae.decode(latents, return_dict=False)[0]
image = self.image_processor.postprocess(image, output_type=output_type)
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (image,)
return FluxPipelineOutput(images=image)

176
toolkit/models/flux.py Normal file
View File

@@ -0,0 +1,176 @@
# forward that bypasses the guidance embedding so it can be avoided during training.
from functools import partial
from typing import Optional
import torch
from diffusers import FluxTransformer2DModel
def guidance_embed_bypass_forward(self, timestep, guidance, pooled_projection):
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(
timesteps_proj.to(dtype=pooled_projection.dtype)) # (N, D)
pooled_projections = self.text_embedder(pooled_projection)
conditioning = timesteps_emb + pooled_projections
return conditioning
# bypass the forward function
def bypass_flux_guidance(transformer):
if hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
return
# dont bypass if it doesnt have the guidance embedding
if not hasattr(transformer.time_text_embed, 'guidance_embedder'):
return
transformer.time_text_embed._bfg_orig_forward = transformer.time_text_embed.forward
transformer.time_text_embed.forward = partial(
guidance_embed_bypass_forward, transformer.time_text_embed
)
# restore the forward function
def restore_flux_guidance(transformer):
if not hasattr(transformer.time_text_embed, '_bfg_orig_forward'):
return
transformer.time_text_embed.forward = transformer.time_text_embed._bfg_orig_forward
del transformer.time_text_embed._bfg_orig_forward
def new_device_to(self: FluxTransformer2DModel, *args, **kwargs):
# Store original device if provided in args or kwargs
device_in_kwargs = 'device' in kwargs
device_in_args = any(isinstance(arg, (str, torch.device)) for arg in args)
device = None
# Remove device from kwargs if present
if device_in_kwargs:
device = kwargs['device']
del kwargs['device']
# Only filter args if we detected a device argument
if device_in_args:
args = list(args)
for idx, arg in enumerate(args):
if isinstance(arg, (str, torch.device)):
device = arg
del args[idx]
self.pos_embed = self.pos_embed.to(device, *args, **kwargs)
self.time_text_embed = self.time_text_embed.to(device, *args, **kwargs)
self.context_embedder = self.context_embedder.to(device, *args, **kwargs)
self.x_embedder = self.x_embedder.to(device, *args, **kwargs)
for block in self.transformer_blocks:
block.to(block._split_device, *args, **kwargs)
for block in self.single_transformer_blocks:
block.to(block._split_device, *args, **kwargs)
self.norm_out = self.norm_out.to(device, *args, **kwargs)
self.proj_out = self.proj_out.to(device, *args, **kwargs)
return self
def split_gpu_double_block_forward(
self,
hidden_states: torch.FloatTensor,
encoder_hidden_states: torch.FloatTensor,
temb: torch.FloatTensor,
image_rotary_emb=None,
joint_attention_kwargs=None,
):
if hidden_states.device != self._split_device:
hidden_states = hidden_states.to(self._split_device)
if encoder_hidden_states.device != self._split_device:
encoder_hidden_states = encoder_hidden_states.to(self._split_device)
if temb.device != self._split_device:
temb = temb.to(self._split_device)
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
# is a tuple of tensors
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
return self._pre_gpu_split_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
def split_gpu_single_block_forward(
self,
hidden_states: torch.FloatTensor,
temb: torch.FloatTensor,
image_rotary_emb=None,
joint_attention_kwargs=None,
**kwargs
):
if hidden_states.device != self._split_device:
hidden_states = hidden_states.to(device=self._split_device)
if temb.device != self._split_device:
temb = temb.to(device=self._split_device)
if image_rotary_emb is not None and image_rotary_emb[0].device != self._split_device:
# is a tuple of tensors
image_rotary_emb = tuple([t.to(self._split_device) for t in image_rotary_emb])
hidden_state_out = self._pre_gpu_split_forward(hidden_states, temb, image_rotary_emb, joint_attention_kwargs, **kwargs)
if hasattr(self, "_split_output_device"):
return hidden_state_out.to(self._split_output_device)
return hidden_state_out
def add_model_gpu_splitter_to_flux(
transformer: FluxTransformer2DModel,
# ~ 5 billion for all other params
other_module_params: Optional[int] = 5e9,
# since they are not trainable, multiply by smaller number
other_module_param_count_scale: Optional[float] = 0.3
):
gpu_id_list = [i for i in range(torch.cuda.device_count())]
# if len(gpu_id_list) > 2:
# raise ValueError("Cannot split to more than 2 GPUs currently.")
other_module_params *= other_module_param_count_scale
# since we are not tuning the
total_params = sum(p.numel() for p in transformer.parameters()) + other_module_params
params_per_gpu = total_params / len(gpu_id_list)
current_gpu_idx = 0
# text encoders, vae, and some non block layers will all be on gpu 0
current_gpu_params = other_module_params
for double_block in transformer.transformer_blocks:
device = torch.device(f"cuda:{current_gpu_idx}")
double_block._pre_gpu_split_forward = double_block.forward
double_block.forward = partial(
split_gpu_double_block_forward, double_block)
double_block._split_device = device
# add the params to the current gpu
current_gpu_params += sum(p.numel() for p in double_block.parameters())
# if the current gpu params are greater than the params per gpu, move to next gpu
if current_gpu_params > params_per_gpu:
current_gpu_idx += 1
current_gpu_params = 0
if current_gpu_idx >= len(gpu_id_list):
current_gpu_idx = gpu_id_list[-1]
for single_block in transformer.single_transformer_blocks:
device = torch.device(f"cuda:{current_gpu_idx}")
single_block._pre_gpu_split_forward = single_block.forward
single_block.forward = partial(
split_gpu_single_block_forward, single_block)
single_block._split_device = device
# add the params to the current gpu
current_gpu_params += sum(p.numel() for p in single_block.parameters())
# if the current gpu params are greater than the params per gpu, move to next gpu
if current_gpu_params > params_per_gpu:
current_gpu_idx += 1
current_gpu_params = 0
if current_gpu_idx >= len(gpu_id_list):
current_gpu_idx = gpu_id_list[-1]
# add output device to last layer
transformer.single_transformer_blocks[-1]._split_output_device = torch.device("cuda:0")
transformer._pre_gpu_split_to = transformer.to
transformer.to = partial(new_device_to, transformer)

View File

@@ -0,0 +1,94 @@
from typing import Optional
from diffusers.models.attention_processor import Attention
import torch
import torch.nn.functional as F
class FluxSageAttnProcessor2_0:
"""Attention processor used typically in processing the SD3-like self-attention projections."""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("FluxAttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn: Attention,
hidden_states: torch.FloatTensor,
encoder_hidden_states: torch.FloatTensor = None,
attention_mask: Optional[torch.FloatTensor] = None,
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.FloatTensor:
from sageattention import sageattn
batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
# `sample` projections.
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
# the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states`
if encoder_hidden_states is not None:
# `context` projections.
encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states)
encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states)
encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states)
encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view(
batch_size, -1, attn.heads, head_dim
).transpose(1, 2)
encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view(
batch_size, -1, attn.heads, head_dim
).transpose(1, 2)
encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view(
batch_size, -1, attn.heads, head_dim
).transpose(1, 2)
if attn.norm_added_q is not None:
encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj)
if attn.norm_added_k is not None:
encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj)
# attention
query = torch.cat([encoder_hidden_states_query_proj, query], dim=2)
key = torch.cat([encoder_hidden_states_key_proj, key], dim=2)
value = torch.cat([encoder_hidden_states_value_proj, value], dim=2)
if image_rotary_emb is not None:
from diffusers.models.embeddings import apply_rotary_emb
query = apply_rotary_emb(query, image_rotary_emb)
key = apply_rotary_emb(key, image_rotary_emb)
hidden_states = sageattn(query, key, value, dropout_p=0.0, is_causal=False)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
if encoder_hidden_states is not None:
encoder_hidden_states, hidden_states = (
hidden_states[:, : encoder_hidden_states.shape[1]],
hidden_states[:, encoder_hidden_states.shape[1] :],
)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
return hidden_states, encoder_hidden_states
else:
return hidden_states

364
toolkit/models/ilora.py Normal file
View File

@@ -0,0 +1,364 @@
import math
import weakref
import torch
import torch.nn as nn
from typing import TYPE_CHECKING, List, Dict, Any
from toolkit.models.clip_fusion import ZipperBlock
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
import sys
from toolkit.paths import REPOS_ROOT
sys.path.append(REPOS_ROOT)
from ipadapter.ip_adapter.resampler import Resampler
from collections import OrderedDict
if TYPE_CHECKING:
from toolkit.lora_special import LoRAModule
from toolkit.stable_diffusion_model import StableDiffusion
class MLP(nn.Module):
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
super().__init__()
if use_residual:
assert in_dim == out_dim
self.layernorm = nn.LayerNorm(in_dim)
self.fc1 = nn.Linear(in_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, out_dim)
self.dropout = nn.Dropout(dropout)
self.use_residual = use_residual
self.act_fn = nn.GELU()
def forward(self, x):
residual = x
x = self.layernorm(x)
x = self.fc1(x)
x = self.act_fn(x)
x = self.fc2(x)
x = self.dropout(x)
if self.use_residual:
x = x + residual
return x
class LoRAGenerator(torch.nn.Module):
def __init__(
self,
input_size: int = 768, # projection dimension
hidden_size: int = 768,
head_size: int = 512,
num_heads: int = 1,
num_mlp_layers: int = 1,
output_size: int = 768,
dropout: float = 0.0
):
super().__init__()
self.input_size = input_size
self.num_heads = num_heads
self.simple = False
self.output_size = output_size
if self.simple:
self.head = nn.Linear(input_size, head_size, bias=False)
else:
self.lin_in = nn.Linear(input_size, hidden_size)
self.mlp_blocks = nn.Sequential(*[
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in range(num_mlp_layers)
])
self.head = nn.Linear(hidden_size, head_size, bias=False)
self.norm = nn.LayerNorm(head_size)
if num_heads == 1:
self.output = nn.Linear(head_size, self.output_size)
# for each output block. multiply weights by 0.01
with torch.no_grad():
self.output.weight.data *= 0.01
else:
head_output_size = output_size // num_heads
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
# for each output block. multiply weights by 0.01
with torch.no_grad():
for output in self.outputs:
output.weight.data *= 0.01
# allow get device
@property
def device(self):
return next(self.parameters()).device
@property
def dtype(self):
return next(self.parameters()).dtype
def forward(self, embedding):
if len(embedding.shape) == 2:
embedding = embedding.unsqueeze(1)
x = embedding
if not self.simple:
x = self.lin_in(embedding)
x = self.mlp_blocks(x)
x = self.head(x)
x = self.norm(x)
if self.num_heads == 1:
x = self.output(x)
else:
out_chunks = torch.chunk(x, self.num_heads, dim=1)
x = []
for out_layer, chunk in zip(self.outputs, out_chunks):
x.append(out_layer(chunk))
x = torch.cat(x, dim=-1)
return x.squeeze(1)
class InstantLoRAMidModule(torch.nn.Module):
def __init__(
self,
index: int,
lora_module: 'LoRAModule',
instant_lora_module: 'InstantLoRAModule',
up_shape: list = None,
down_shape: list = None,
):
super(InstantLoRAMidModule, self).__init__()
self.up_shape = up_shape
self.down_shape = down_shape
self.index = index
self.lora_module_ref = weakref.ref(lora_module)
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
self.embed = None
def down_forward(self, x, *args, **kwargs):
# get the embed
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
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]
batch_size = x.shape[0]
# unconditional
if down_weight.shape[0] * 2 == batch_size:
down_weight = torch.cat([down_weight] * 2, dim=0)
weight_chunks = torch.chunk(down_weight, batch_size, dim=0)
x_chunks = torch.chunk(x, batch_size, dim=0)
x_out = []
for i in range(batch_size):
weight_chunk = weight_chunks[i]
x_chunk = x_chunks[i]
# reshape
weight_chunk = weight_chunk.view(self.down_shape)
# check if is conv or linear
if len(weight_chunk.shape) == 4:
org_module = self.lora_module_ref().orig_module_ref()
stride = org_module.stride
padding = org_module.padding
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding, stride=stride)
else:
# run a simple linear layer with the down weight
x_chunk = x_chunk @ weight_chunk.T
x_out.append(x_chunk)
x = torch.cat(x_out, dim=0)
return x
def up_forward(self, x, *args, **kwargs):
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
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:]
batch_size = x.shape[0]
# unconditional
if up_weight.shape[0] * 2 == batch_size:
up_weight = torch.cat([up_weight] * 2, dim=0)
weight_chunks = torch.chunk(up_weight, batch_size, dim=0)
x_chunks = torch.chunk(x, batch_size, dim=0)
x_out = []
for i in range(batch_size):
weight_chunk = weight_chunks[i]
x_chunk = x_chunks[i]
# reshape
weight_chunk = weight_chunk.view(self.up_shape)
# check if is conv or linear
if len(weight_chunk.shape) == 4:
padding = 0
if weight_chunk.shape[-1] == 3:
padding = 1
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
else:
# run a simple linear layer with the down weight
x_chunk = x_chunk @ weight_chunk.T
x_out.append(x_chunk)
x = torch.cat(x_out, dim=0)
return x
class InstantLoRAModule(torch.nn.Module):
def __init__(
self,
vision_hidden_size: int,
vision_tokens: int,
head_dim: int,
num_heads: int, # number of heads in the resampler
sd: 'StableDiffusion',
config=None
):
super(InstantLoRAModule, self).__init__()
# self.linear = torch.nn.Linear(2, 1)
self.sd_ref = weakref.ref(sd)
self.dim = sd.network.lora_dim
self.vision_hidden_size = vision_hidden_size
self.vision_tokens = vision_tokens
self.head_dim = head_dim
self.num_heads = num_heads
# stores the projection vector. Grabbed by modules
self.img_embeds: List[torch.Tensor] = None
# disable merging in. It is slower on inference
self.sd_ref().network.can_merge_in = False
self.ilora_modules = torch.nn.ModuleList()
lora_modules = self.sd_ref().network.get_all_modules()
output_size = 0
self.embed_lengths = []
self.weight_mapping = []
for idx, lora_module in enumerate(lora_modules):
module_dict = lora_module.state_dict()
down_shape = list(module_dict['lora_down.weight'].shape)
up_shape = list(module_dict['lora_up.weight'].shape)
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
module_size = math.prod(down_shape) + math.prod(up_shape)
output_size += module_size
self.embed_lengths.append(module_size)
# add a new mid module that will take the original forward and add a vector to it
# this will be used to add the vector to the original forward
instant_module = InstantLoRAMidModule(
idx,
lora_module,
self,
up_shape=up_shape,
down_shape=down_shape
)
self.ilora_modules.append(instant_module)
# replace the LoRA forwards
lora_module.lora_down.forward = instant_module.down_forward
lora_module.lora_up.forward = instant_module.up_forward
self.output_size = output_size
number_formatted_output_size = "{:,}".format(output_size)
print(f" ILORA output size: {number_formatted_output_size}")
# if not evenly divisible, error
if self.output_size % self.num_heads != 0:
raise ValueError("Output size must be divisible by the number of heads")
self.head_output_size = self.output_size // self.num_heads
if vision_tokens > 1:
self.resampler = Resampler(
dim=vision_hidden_size,
depth=4,
dim_head=64,
heads=12,
num_queries=num_heads, # output tokens
embedding_dim=vision_hidden_size,
max_seq_len=vision_tokens,
output_dim=head_dim,
apply_pos_emb=True, # this is new
ff_mult=4
)
self.proj_module = LoRAGenerator(
input_size=head_dim,
hidden_size=head_dim,
head_size=head_dim,
num_mlp_layers=1,
num_heads=self.num_heads,
output_size=self.output_size,
)
self.migrate_weight_mapping()
def migrate_weight_mapping(self):
return
# # changes the names of the modules to common ones
# keymap = self.sd_ref().network.get_keymap()
# save_keymap = {}
# if keymap is not None:
# for ldm_key, diffusers_key in keymap.items():
# # invert them
# save_keymap[diffusers_key] = ldm_key
#
# new_keymap = {}
# for key, value in self.weight_mapping:
# if key in save_keymap:
# new_keymap[save_keymap[key]] = value
# else:
# print(f"Key {key} not found in keymap")
# new_keymap[key] = value
# self.weight_mapping = new_keymap
# else:
# print("No keymap found. Using default names")
# return
def forward(self, img_embeds):
# expand token rank if only rank 2
if len(img_embeds.shape) == 2:
img_embeds = img_embeds.unsqueeze(1)
# resample the image embeddings
img_embeds = self.resampler(img_embeds)
img_embeds = self.proj_module(img_embeds)
if len(img_embeds.shape) == 3:
# merge the heads
img_embeds = img_embeds.mean(dim=1)
self.img_embeds = []
# get all the slices
start = 0
for length in self.embed_lengths:
self.img_embeds.append(img_embeds[:, start:start+length])
start += length
def get_additional_save_metadata(self) -> Dict[str, Any]:
# save the weight mapping
return {
"weight_mapping": self.weight_mapping,
"num_heads": self.num_heads,
"vision_hidden_size": self.vision_hidden_size,
"head_dim": self.head_dim,
"vision_tokens": self.vision_tokens,
"output_size": self.output_size,
}

419
toolkit/models/ilora2.py Normal file
View File

@@ -0,0 +1,419 @@
import math
import weakref
from toolkit.config_modules import AdapterConfig
import torch
import torch.nn as nn
from typing import TYPE_CHECKING, List, Dict, Any
from toolkit.models.clip_fusion import ZipperBlock
from toolkit.models.zipper_resampler import ZipperModule, ZipperResampler
import sys
from toolkit.paths import REPOS_ROOT
sys.path.append(REPOS_ROOT)
from ipadapter.ip_adapter.resampler import Resampler
from collections import OrderedDict
if TYPE_CHECKING:
from toolkit.lora_special import LoRAModule
from toolkit.stable_diffusion_model import StableDiffusion
class MLP(nn.Module):
def __init__(self, in_dim, out_dim, hidden_dim, dropout=0.1, use_residual=True):
super().__init__()
if use_residual:
assert in_dim == out_dim
self.layernorm = nn.LayerNorm(in_dim)
self.fc1 = nn.Linear(in_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, out_dim)
self.dropout = nn.Dropout(dropout)
self.use_residual = use_residual
self.act_fn = nn.GELU()
def forward(self, x):
residual = x
x = self.layernorm(x)
x = self.fc1(x)
x = self.act_fn(x)
x = self.fc2(x)
x = self.dropout(x)
if self.use_residual:
x = x + residual
return x
class LoRAGenerator(torch.nn.Module):
def __init__(
self,
input_size: int = 768, # projection dimension
hidden_size: int = 768,
head_size: int = 512,
num_heads: int = 1,
num_mlp_layers: int = 1,
output_size: int = 768,
dropout: float = 0.0
):
super().__init__()
self.input_size = input_size
self.num_heads = num_heads
self.simple = False
self.output_size = output_size
if self.simple:
self.head = nn.Linear(input_size, head_size, bias=False)
else:
self.lin_in = nn.Linear(input_size, hidden_size)
self.mlp_blocks = nn.Sequential(*[
MLP(hidden_size, hidden_size, hidden_size, dropout=dropout, use_residual=True) for _ in
range(num_mlp_layers)
])
self.head = nn.Linear(hidden_size, head_size, bias=False)
self.norm = nn.LayerNorm(head_size)
if num_heads == 1:
self.output = nn.Linear(head_size, self.output_size)
# for each output block. multiply weights by 0.01
with torch.no_grad():
self.output.weight.data *= 0.01
else:
head_output_size = output_size // num_heads
self.outputs = nn.ModuleList([nn.Linear(head_size, head_output_size) for _ in range(num_heads)])
# for each output block. multiply weights by 0.01
with torch.no_grad():
for output in self.outputs:
output.weight.data *= 0.01
# allow get device
@property
def device(self):
return next(self.parameters()).device
@property
def dtype(self):
return next(self.parameters()).dtype
def forward(self, embedding):
if len(embedding.shape) == 2:
embedding = embedding.unsqueeze(1)
x = embedding
if not self.simple:
x = self.lin_in(embedding)
x = self.mlp_blocks(x)
x = self.head(x)
x = self.norm(x)
if self.num_heads == 1:
x = self.output(x)
else:
out_chunks = torch.chunk(x, self.num_heads, dim=1)
x = []
for out_layer, chunk in zip(self.outputs, out_chunks):
x.append(out_layer(chunk))
x = torch.cat(x, dim=-1)
return x.squeeze(1)
class InstantLoRAMidModule(torch.nn.Module):
def __init__(
self,
index: int,
lora_module: 'LoRAModule',
instant_lora_module: 'InstantLoRAModule',
up_shape: list = None,
down_shape: list = None,
):
super(InstantLoRAMidModule, self).__init__()
self.up_shape = up_shape
self.down_shape = down_shape
self.index = index
self.lora_module_ref = weakref.ref(lora_module)
self.instant_lora_module_ref = weakref.ref(instant_lora_module)
self.do_up = instant_lora_module.config.ilora_up
self.do_down = instant_lora_module.config.ilora_down
self.do_mid = instant_lora_module.config.ilora_mid
self.down_dim = self.down_shape[1] if self.do_down else 0
self.mid_dim = self.up_shape[1] if self.do_mid else 0
self.out_dim = self.up_shape[0] if self.do_up else 0
self.embed = None
def down_forward(self, x, *args, **kwargs):
if not self.do_down:
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
# get the embed
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
down_weight = self.embed[:, :self.down_dim]
batch_size = x.shape[0]
# unconditional
if down_weight.shape[0] * 2 == batch_size:
down_weight = torch.cat([down_weight] * 2, dim=0)
try:
if len(x.shape) == 4:
# conv
down_weight = down_weight.view(batch_size, -1, 1, 1)
if x.shape[1] != down_weight.shape[1]:
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
elif len(x.shape) == 2:
down_weight = down_weight.view(batch_size, -1)
if x.shape[1] != down_weight.shape[1]:
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
else:
down_weight = down_weight.view(batch_size, 1, -1)
if x.shape[2] != down_weight.shape[2]:
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
x = x * down_weight
x = self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
except Exception as e:
print(e)
raise ValueError(f"Down weight shape not understood: {down_weight.shape} {x.shape}")
return x
def up_forward(self, x, *args, **kwargs):
# do mid here
x = self.mid_forward(x, *args, **kwargs)
if not self.do_up:
return self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
# get the embed
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
up_weight = self.embed[:, -self.out_dim:]
batch_size = x.shape[0]
# unconditional
if up_weight.shape[0] * 2 == batch_size:
up_weight = torch.cat([up_weight] * 2, dim=0)
try:
if len(x.shape) == 4:
# conv
up_weight = up_weight.view(batch_size, -1, 1, 1)
elif len(x.shape) == 2:
up_weight = up_weight.view(batch_size, -1)
else:
up_weight = up_weight.view(batch_size, 1, -1)
x = self.lora_module_ref().lora_up.orig_forward(x, *args, **kwargs)
x = x * up_weight
except Exception as e:
print(e)
raise ValueError(f"Up weight shape not understood: {up_weight.shape} {x.shape}")
return x
def mid_forward(self, x, *args, **kwargs):
if not self.do_mid:
return self.lora_module_ref().lora_down.orig_forward(x, *args, **kwargs)
batch_size = x.shape[0]
# get the embed
self.embed = self.instant_lora_module_ref().img_embeds[self.index]
mid_weight = self.embed[:, self.down_dim:self.down_dim + self.mid_dim * self.mid_dim]
# unconditional
if mid_weight.shape[0] * 2 == batch_size:
mid_weight = torch.cat([mid_weight] * 2, dim=0)
weight_chunks = torch.chunk(mid_weight, batch_size, dim=0)
x_chunks = torch.chunk(x, batch_size, dim=0)
x_out = []
for i in range(batch_size):
weight_chunk = weight_chunks[i]
x_chunk = x_chunks[i]
# reshape
if len(x_chunk.shape) == 4:
# conv
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim, 1, 1)
else:
weight_chunk = weight_chunk.view(self.mid_dim, self.mid_dim)
# check if is conv or linear
if len(weight_chunk.shape) == 4:
padding = 0
if weight_chunk.shape[-1] == 3:
padding = 1
x_chunk = nn.functional.conv2d(x_chunk, weight_chunk, padding=padding)
else:
# run a simple linear layer with the down weight
x_chunk = x_chunk @ weight_chunk.T
x_out.append(x_chunk)
x = torch.cat(x_out, dim=0)
return x
class InstantLoRAModule(torch.nn.Module):
def __init__(
self,
vision_hidden_size: int,
vision_tokens: int,
head_dim: int,
num_heads: int, # number of heads in the resampler
sd: 'StableDiffusion',
config: AdapterConfig
):
super(InstantLoRAModule, self).__init__()
# self.linear = torch.nn.Linear(2, 1)
self.sd_ref = weakref.ref(sd)
self.dim = sd.network.lora_dim
self.vision_hidden_size = vision_hidden_size
self.vision_tokens = vision_tokens
self.head_dim = head_dim
self.num_heads = num_heads
self.config: AdapterConfig = config
# stores the projection vector. Grabbed by modules
self.img_embeds: List[torch.Tensor] = None
# disable merging in. It is slower on inference
self.sd_ref().network.can_merge_in = False
self.ilora_modules = torch.nn.ModuleList()
lora_modules = self.sd_ref().network.get_all_modules()
output_size = 0
self.embed_lengths = []
self.weight_mapping = []
for idx, lora_module in enumerate(lora_modules):
module_dict = lora_module.state_dict()
down_shape = list(module_dict['lora_down.weight'].shape)
up_shape = list(module_dict['lora_up.weight'].shape)
self.weight_mapping.append([lora_module.lora_name, [down_shape, up_shape]])
#
# module_size = math.prod(down_shape) + math.prod(up_shape)
# conv weight shape is (out_channels, in_channels, kernel_size, kernel_size)
# linear weight shape is (out_features, in_features)
# just doing in dim and out dim
in_dim = down_shape[1] if self.config.ilora_down else 0
mid_dim = down_shape[0] * down_shape[0] if self.config.ilora_mid else 0
out_dim = up_shape[0] if self.config.ilora_up else 0
module_size = in_dim + mid_dim + out_dim
output_size += module_size
self.embed_lengths.append(module_size)
# add a new mid module that will take the original forward and add a vector to it
# this will be used to add the vector to the original forward
instant_module = InstantLoRAMidModule(
idx,
lora_module,
self,
up_shape=up_shape,
down_shape=down_shape
)
self.ilora_modules.append(instant_module)
# replace the LoRA forwards
lora_module.lora_down.orig_forward = lora_module.lora_down.forward
lora_module.lora_down.forward = instant_module.down_forward
lora_module.lora_up.orig_forward = lora_module.lora_up.forward
lora_module.lora_up.forward = instant_module.up_forward
self.output_size = output_size
number_formatted_output_size = "{:,}".format(output_size)
print(f" ILORA output size: {number_formatted_output_size}")
# if not evenly divisible, error
if self.output_size % self.num_heads != 0:
raise ValueError("Output size must be divisible by the number of heads")
self.head_output_size = self.output_size // self.num_heads
if vision_tokens > 1:
self.resampler = Resampler(
dim=vision_hidden_size,
depth=4,
dim_head=64,
heads=12,
num_queries=num_heads, # output tokens
embedding_dim=vision_hidden_size,
max_seq_len=vision_tokens,
output_dim=head_dim,
apply_pos_emb=True, # this is new
ff_mult=4
)
self.proj_module = LoRAGenerator(
input_size=head_dim,
hidden_size=head_dim,
head_size=head_dim,
num_mlp_layers=1,
num_heads=self.num_heads,
output_size=self.output_size,
)
self.migrate_weight_mapping()
def migrate_weight_mapping(self):
return
# # changes the names of the modules to common ones
# keymap = self.sd_ref().network.get_keymap()
# save_keymap = {}
# if keymap is not None:
# for ldm_key, diffusers_key in keymap.items():
# # invert them
# save_keymap[diffusers_key] = ldm_key
#
# new_keymap = {}
# for key, value in self.weight_mapping:
# if key in save_keymap:
# new_keymap[save_keymap[key]] = value
# else:
# print(f"Key {key} not found in keymap")
# new_keymap[key] = value
# self.weight_mapping = new_keymap
# else:
# print("No keymap found. Using default names")
# return
def forward(self, img_embeds):
# expand token rank if only rank 2
if len(img_embeds.shape) == 2:
img_embeds = img_embeds.unsqueeze(1)
# resample the image embeddings
img_embeds = self.resampler(img_embeds)
img_embeds = self.proj_module(img_embeds)
if len(img_embeds.shape) == 3:
# merge the heads
img_embeds = img_embeds.mean(dim=1)
self.img_embeds = []
# get all the slices
start = 0
for length in self.embed_lengths:
self.img_embeds.append(img_embeds[:, start:start + length])
start += length
def get_additional_save_metadata(self) -> Dict[str, Any]:
# save the weight mapping
return {
"weight_mapping": self.weight_mapping,
"num_heads": self.num_heads,
"vision_hidden_size": self.vision_hidden_size,
"head_dim": self.head_dim,
"vision_tokens": self.vision_tokens,
"output_size": self.output_size,
"do_up": self.config.ilora_up,
"do_mid": self.config.ilora_mid,
"do_down": self.config.ilora_down,
}

View File

@@ -0,0 +1,191 @@
from functools import partial
import sys
import torch
import torch.nn as nn
import torch.nn.functional as F
import weakref
from typing import Any, Dict, List, Optional, Tuple, Union, TYPE_CHECKING
from diffusers.models.transformers.transformer_flux import FluxTransformerBlock
from transformers import AutoModel, AutoTokenizer, Qwen2Model, LlamaModel, Qwen2Tokenizer, LlamaTokenizer
from toolkit import train_tools
from toolkit.prompt_utils import PromptEmbeds
from diffusers import Transformer2DModel
from toolkit.dequantize import patch_dequantization_on_save
if TYPE_CHECKING:
from toolkit.stable_diffusion_model import StableDiffusion, PixArtSigmaPipeline
from toolkit.custom_adapter import CustomAdapter
LLM = Union[Qwen2Model, LlamaModel]
LLMTokenizer = Union[Qwen2Tokenizer, LlamaTokenizer]
def new_context_embedder_forward(self, x):
if self._adapter_ref().is_active:
x = self._context_embedder_ref()(x)
else:
x = self._orig_forward(x)
return x
def new_block_forward(
self: FluxTransformerBlock,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
if self._adapter_ref().is_active:
return self._new_block_ref()(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
else:
return self._orig_forward(hidden_states, encoder_hidden_states, temb, image_rotary_emb, joint_attention_kwargs)
class LLMAdapter(torch.nn.Module):
def __init__(
self,
adapter: 'CustomAdapter',
sd: 'StableDiffusion',
llm: LLM,
tokenizer: LLMTokenizer,
num_cloned_blocks: int = 0,
):
super(LLMAdapter, self).__init__()
self.adapter_ref: weakref.ref = weakref.ref(adapter)
self.sd_ref: weakref.ref = weakref.ref(sd)
self.llm_ref: weakref.ref = weakref.ref(llm)
self.tokenizer_ref: weakref.ref = weakref.ref(tokenizer)
self.num_cloned_blocks = num_cloned_blocks
self.apply_embedding_mask = False
# make sure we can pad
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
# self.system_prompt = ""
self.system_prompt = "You are an assistant designed to generate superior images with the superior degree of image-text alignment based on textual prompts or user prompts. <Prompt Start> "
# determine length of system prompt
sys_prompt_tokenized = tokenizer(
[self.system_prompt],
padding="longest",
return_tensors="pt",
)
sys_prompt_tokenized_ids = sys_prompt_tokenized.input_ids[0]
self.system_prompt_length = sys_prompt_tokenized_ids.shape[0]
print(f"System prompt length: {self.system_prompt_length}")
self.hidden_size = llm.config.hidden_size
blocks = []
if sd.is_flux:
self.apply_embedding_mask = True
self.context_embedder = nn.Linear(
self.hidden_size, sd.unet.inner_dim)
self.sequence_length = 512
sd.unet.context_embedder._orig_forward = sd.unet.context_embedder.forward
sd.unet.context_embedder.forward = partial(
new_context_embedder_forward, sd.unet.context_embedder)
sd.unet.context_embedder._context_embedder_ref = weakref.ref(self.context_embedder)
# add a is active property to the context embedder
sd.unet.context_embedder._adapter_ref = self.adapter_ref
for idx in range(self.num_cloned_blocks):
block = FluxTransformerBlock(
dim=sd.unet.inner_dim,
num_attention_heads=24,
attention_head_dim=128,
)
# patch it in case it is quantized
patch_dequantization_on_save(sd.unet.transformer_blocks[idx])
state_dict = sd.unet.transformer_blocks[idx].state_dict()
for key, value in state_dict.items():
block.state_dict()[key].copy_(value)
blocks.append(block)
orig_block = sd.unet.transformer_blocks[idx]
orig_block._orig_forward = orig_block.forward
orig_block.forward = partial(
new_block_forward, orig_block)
orig_block._new_block_ref = weakref.ref(block)
orig_block._adapter_ref = self.adapter_ref
elif sd.is_lumina2:
self.context_embedder = nn.Linear(
self.hidden_size, sd.unet.hidden_size)
self.sequence_length = 256
else:
raise ValueError(
"llm adapter currently only supports flux or lumina2")
self.blocks = nn.ModuleList(blocks)
def _get_prompt_embeds(
self,
prompt: Union[str, List[str]],
max_sequence_length: int = 256,
) -> Tuple[torch.Tensor, torch.Tensor]:
tokenizer = self.tokenizer_ref()
text_encoder = self.llm_ref()
device = text_encoder.device
prompt = [prompt] if isinstance(prompt, str) else prompt
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length + self.system_prompt_length,
truncation=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids.to(device)
prompt_attention_mask = text_inputs.attention_mask.to(device)
# remove the system prompt from the input and attention mask
prompt_embeds = text_encoder(
text_input_ids, attention_mask=prompt_attention_mask, output_hidden_states=True
)
prompt_embeds = prompt_embeds.hidden_states[-1]
prompt_embeds = prompt_embeds[:, self.system_prompt_length:]
prompt_attention_mask = prompt_attention_mask[:, self.system_prompt_length:]
dtype = text_encoder.dtype
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
return prompt_embeds, prompt_attention_mask
# make a getter to see if is active
@property
def is_active(self):
return self.adapter_ref().is_active
def encode_text(self, prompt):
prompt = prompt if isinstance(prompt, list) else [prompt]
prompt = [self.system_prompt + p for p in prompt]
# prompt = [self.system_prompt + p for p in prompt]
prompt_embeds, prompt_attention_mask = self._get_prompt_embeds(
prompt=prompt,
max_sequence_length=self.sequence_length,
)
prompt_embeds = PromptEmbeds(
prompt_embeds,
attention_mask=prompt_attention_mask,
).detach()
return prompt_embeds
def forward(self, input):
return input

Some files were not shown because too many files have changed in this diff Show More