370 Commits

Author SHA1 Message Date
Jaret Burkett
4fa8fac5fd WIP multidevice training 2024-08-29 16:04:20 -06:00
Jaret Burkett
a48c9aba8d Created a v2 trainer and moved all the training logic to single torch model so it can can be run in parallel 2024-08-29 12:34:18 -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
c2a4b8e058 More timestep fixes 2023-11-28 10:34:08 -07:00
Jaret Burkett
792a5e37e2 Numerous fixes for time sampling. Still not perfect 2023-11-28 07:34:43 -07:00
Jaret Burkett
d7e55b6ad4 Bug fixes, negative prompting during training, hardened catching 2023-11-24 07:25:11 -07:00
Jaret Burkett
fbec68681d Added timestep modifications to lcm scheduler for more evenly spaced timesteps 2023-11-17 23:26:52 -07:00
Jaret Burkett
6280284d8b Fixed cleanup of emebddings. 2023-11-16 20:26:11 -07:00
Jaret Burkett
ad50921c41 Sampling tests and added fixes for cleanups 2023-11-16 08:33:23 -07:00
Jaret Burkett
e47006ed70 Added some features for an LCM condenser plugin 2023-11-15 08:56:45 -07:00
Jaret Burkett
4f9cdd916a added prompt dropout to happen indempendently on each TE 2023-11-14 05:26:51 -07:00
Jaret Burkett
7782caa468 Varous bug fixes. Finalized targeted guidance algo 2023-11-10 12:18:08 -07:00
Jaret Burkett
fa6d91ba76 Diffirential guidance working, but I may have a better way 2023-11-09 06:09:36 -07:00
Jaret Burkett
1ee62562a4 diffirential guidance is WORKING (from what I can tell) 2023-11-07 19:24:12 -07:00
Jaret Burkett
dc8448d958 Added a way pass refiner ratio to sample config 2023-11-06 09:22:58 -07:00
Jaret Burkett
a8b3b8b8da Fixed weight mapping for refiner 2023-11-06 07:37:47 -07:00
Jaret Burkett
93ea955d7c Added refiner fine tuning. Works, but needs some polish. 2023-11-05 17:15:03 -07:00
Jaret Burkett
8a9e8f708f Added base for using guidance during training. Still not working right. 2023-11-05 04:03:32 -07:00
Jaret Burkett
d35733ac06 Added support for training ssd-1B. Added support for saving models into diffusers format. We can currently save in safetensors format for ssd-1b, but diffusers cannot load it yet. 2023-11-03 05:01:16 -06:00
Jaret Burkett
ceaf1d9454 Various bug fixes, wip stuff, and tweaks 2023-11-02 18:19:20 -06:00
Jaret Burkett
7d707b2fe6 Added masking to slider training. Something is still weird though 2023-11-01 14:51:29 -06:00
Jaret Burkett
a899ec91c8 Added some split prompting started code, adamw8bit, replacements improving, learnable snr gos. A lot of good stuff. 2023-11-01 06:52:21 -06:00
Jaret Burkett
436a09430e Added flat snr gamma vs min. Fixes timestep timing 2023-10-29 15:41:55 -06:00
Jaret Burkett
3097865203 Fix llava import.. again 2023-10-29 12:52:17 -06:00
Jaret Burkett
b84e3260cb Fix llava import 2023-10-29 12:50:54 -06:00
Jaret Burkett
48a9bac22d Added doffusers schedulers 2023-10-29 12:39:50 -06:00
Jaret Burkett
298001439a Added gradient accumulation finally 2023-10-28 13:14:29 -06:00
Jaret Burkett
6f3e0d5af2 Improved lorm extraction and training 2023-10-28 08:21:59 -06:00
Jaret Burkett
0a79ac9604 Added lorm. WIP 2023-10-26 18:23:51 -06:00
Jaret Burkett
9636194c09 Added fuyu captioning 2023-10-25 14:14:53 -06:00
Jaret Burkett
d742792ee4 Fixed issue with not loading short prompt 2023-10-24 16:19:32 -06:00
Jaret Burkett
002279cec3 Allow short and long caption combinations like form the new captioning system. Merge the network into the model before inference and reextract when done. Doubles inference speed on locon models during inference. allow splitting a batch into individual components and run them through alone. Basicallt gradient accumulation with single batch size. 2023-10-24 16:02:07 -06:00
Jaret Burkett
73c8b50975 Added ability to use adagrad from transformers 2023-10-24 11:16:01 -06:00
Jaret Burkett
34eb563d55 Added dataset tagging and management tools using llava 2023-10-24 10:11:47 -06:00
Jaret Burkett
dc36bbb3c8 Added long prompts to general training 2023-10-23 08:12:58 -06:00
Jaret Burkett
9905a1e205 Fixes and longer prompts 2023-10-22 08:57:37 -06:00
Jaret Burkett
0e9fc42816 Allow for inverted masked prior 2023-10-21 06:50:17 -06:00
Jaret Burkett
d46112a354 Fixes for aug pipeline 2023-10-20 13:24:07 -06:00
Jaret Burkett
07bf7bd7de Allow augmentations and targeting different loss types fron the config file 2023-10-18 03:04:57 -06:00
Jaret Burkett
da6302ada8 added a method to apply multipliers to noise and latents prior to combining 2023-10-17 06:09:16 -06:00
Jaret Burkett
a05459afaf Fixed issue with adapters that only had 1 input channel. Added ability to set the percentage chance of adapter matching 2023-10-15 15:13:35 -06:00
Jaret Burkett
b1a22d0b3e hardened reading prompts from json 2023-10-15 07:20:33 -06:00
Jaret Burkett
7909b50d24 Added adapter assistance to SD training 2023-10-14 08:44:53 -06:00
Jaret Burkett
38e441a29c allow flipping for point of interesting autocropping. allow num repeats. Fixed some bugs with new free u 2023-10-12 21:02:47 -06:00
Jaret Burkett
4e3b2c2569 Allow for alpha to be used as a mask 2023-10-10 17:08:59 -06:00
Jaret Burkett
239addba51 Fixed memory leak when cachine latents to disk 2023-10-10 13:31:47 -06:00
Jaret Burkett
63ceffae24 Massive speed increases and ram optimizations 2023-10-10 06:07:55 -06:00
Jaret Burkett
f4c90bb589 Added config to set the min value of a mask 2023-10-09 15:47:54 -06:00
Jaret Burkett
bb1d3793e3 Added ability to add masks to dataloader and sd trainer to adjust weight of image 2023-10-09 11:21:00 -06:00
Jaret Burkett
1d3de678aa fixed bug with trigger word embedding. Allow control images to load from the dataloader or legacy way 2023-10-09 06:21:49 -06:00
Jaret Burkett
b1cfafa0c6 always prompt both even when training only one encoder 2023-10-07 11:23:09 -06:00
Jaret Burkett
cac8754399 Allow for training loras on onle one text encoder for sdxl 2023-10-06 08:11:56 -06:00
Jaret Burkett
f73402473b Bug fixes. Added some functionality to help with private extensions 2023-10-05 07:09:34 -06:00
Jaret Burkett
579650eaf8 Fixed big issue with bucketing dataloader and added random cripping to a point of interest 2023-10-02 18:31:08 -06:00
Jaret Burkett
320e109c5f Allow loading from a json detail file for captions 2023-10-02 13:25:09 -06:00
Jaret Burkett
560251a24f fixed issue with down block residuals when doing slider cfg on sdxl with t2i adapter assisted training 2023-10-01 07:32:48 -06:00
Jaret Burkett
085787b799 Allow loading auxillery images from dataloader 2023-09-30 07:28:23 -06:00
Jaret Burkett
8d9450ad7c Compatability fixes 2023-09-29 14:07:37 -06:00
Jaret Burkett
8509da60cb Added a way to add a t2i adapter guided slider training for more consitant images 2023-09-28 14:08:56 -06:00
Jaret Burkett
c5d49ba661 Allow adapter image to be cropped to match bucket cropping 2023-09-25 10:38:54 -06:00
Jaret Burkett
76c764af49 Fixed issues with adapter and disabled telementry for now 2023-09-24 12:51:29 -06:00
Jaret Burkett
abf7cd221d allow setting adapter weight in prompts 2023-09-24 06:51:54 -06:00
Jaret Burkett
e5153d87c9 Fixed issues with dataloader bucketing. Allow using standard base image for t2i adapters. 2023-09-24 05:19:57 -06:00
Jaret Burkett
830e87cb87 Added IP adapter training. Not functioning correctly yet 2023-09-24 02:39:43 -06:00
Jaret Burkett
19255cdc7c Bugfixes. Added small augmentations to dataloader. Will switch to abluminations soon though. Added ability to adjust step count on start to override what is in the file 2023-09-20 05:30:10 -06:00
Jaret Burkett
0f105690cc Added some further extendability for plugins 2023-09-19 05:41:44 -06:00
Jaret Burkett
61badf85a7 t2i training working from what I can tell at least 2023-09-17 15:56:43 -06:00
Jaret Burkett
181f237a7b added flipping x and y for dataset loader 2023-09-17 08:42:54 -06:00
Jaret Burkett
c698837241 Fixes to esrgan trainer. Moved logic for sd prompt embeddings out of diffusers pipeline so I can manipulate it 2023-09-16 17:41:07 -06:00
Jaret Burkett
27f343fc08 Added base setup for training t2i adapters. Currently untested, saw something else shiny i wanted to finish sirst. Added content_or_style to the training config. It defaults to balanced, which is standard uniform time step sampling. If style or content is passed, it will use cubic sampling for timesteps to favor timesteps that are beneficial for training them. for style, favor later timesteps. For content, favor earlier timesteps. 2023-09-16 08:30:38 -06:00
Jaret Burkett
3eb3535683 Merge pull request #12 from bendeguzvaradi/main
Bug/Safety checker to None
2023-09-14 15:31:09 -06:00
Jaret Burkett
17e4fe40d7 Prevent lycoris network moduels if not training that part of network. Skew timesteps to favor later steps. It performs better 2023-09-14 15:13:24 -06:00
Jaret Burkett
569d7464d5 implemented device placement preset system more places. Vastly improved speed on setting network multiplier and activating network. Fixed timing issues on progress bar 2023-09-14 08:31:54 -06:00
Jaret Burkett
4e945917df added dropout to LoRA networks 2023-09-13 15:23:07 -06:00
Jaret Burkett
ae70200d3c Bug fixes, speed improvements, compatability adjustments withdiffusers updates 2023-09-13 07:03:53 -06:00
Jaret Burkett
d8d1e6fd1e big fixes 2023-09-12 18:48:39 -06:00
Jaret Burkett
257da9493d Dont load image if we are cachine latents 2023-09-12 18:39:41 -06:00
Jaret Burkett
b5a2669b74 Fixed memory leak 2023-09-12 07:03:10 -06:00
Jaret Burkett
d74dd636ee Memory optimizations. Default to using cudamalloc when torch 2.0 for mem allocation 2023-09-12 04:30:23 -06:00
Jaret Burkett
e8583860ad Upgraded to dev for t2i on diffusers. Minor migrations to make it work. 2023-09-11 14:46:06 -06:00
bendeguzvaradi
3d387103cd safety checker to None 2023-09-11 17:05:55 +02:00
Jaret Burkett
083cefa78c Bugfixes for slider reference 2023-09-10 18:36:23 -06:00
Jaret Burkett
b5ec8e4eb1 Improve reference slider memory and speed 2023-09-10 18:26:44 -06:00
Jaret Burkett
708b07adb7 Fixed issue with interleaving when doing cfg 2023-09-10 10:26:58 -06:00
Jaret Burkett
a437aed45f bug fix 2023-09-10 09:52:14 -06:00
Jaret Burkett
34bfeba229 Massive speed increase. Added latent caching both to disk and to memory 2023-09-10 08:54:49 -06:00
Jaret Burkett
41a3f63b72 allow smaller images in buckets and bucket them 2023-09-10 03:43:02 -06:00
Jaret Burkett
626ed2939a bug fixes 2023-09-09 15:04:44 -06:00
Jaret Burkett
2128ac1e08 fixed issue with embed name, save whole config to dir instead of just process so it can be easily shared. Only make one config, no timesteps 2023-09-09 12:24:08 -06:00
Jaret Burkett
be804c9cf5 Save embeddings as their trigger to match auto and comfy style loading. Also, FINALLY found why gradients were wonkey and fixed it. The root problem is dropping out of network state before backward pass. 2023-09-09 12:02:07 -06:00
Jaret Burkett
408c50ead1 actually got gradient checkpointing working, again, again, maybe 2023-09-09 11:27:42 -06:00
Jaret Burkett
4ed03a8d92 Fixed issue with buckets scaling. again 2023-09-08 16:32:14 -06:00
Jaret Burkett
b01ab5d375 FINALLY fixed gradient checkpointing issue. Big batches baby. 2023-09-08 15:21:46 -06:00
Jaret Burkett
cb91b0d6da Changed model download from HF to fp16 2023-09-08 07:57:19 -06:00
Jaret Burkett
ce4f9fe02a Bug fixes and improvements to token injection 2023-09-08 06:10:59 -06:00
Jaret Burkett
92a086d5a5 Fixed issue with token replacements 2023-09-07 13:42:39 -06:00
Jaret Burkett
3feb663a51 Switched to new bucket system that matched sdxl trained buckets. Fixed requirements. Updated embeddings to work with sdxl. Added method to train lora with an embedding at the trigger. Still testing but works amazingly well from what I can see 2023-09-07 13:06:18 -06:00
Jaret Burkett
436bf0c6a3 Added experimental concept replacer, replicate converter, bucket maker, and other goodies 2023-09-06 18:50:32 -06:00
Jaret Burkett
f84500159c Fixed issue with lora layer check 2023-09-04 14:27:37 -06:00
Jaret Burkett
64a5441832 Fully tested and now supporting locon on sdxl. If you have the ram 2023-09-04 14:05:10 -06:00
Jaret Burkett
a4c3507a62 Added LoCON from LyCORIS 2023-09-04 08:48:07 -06:00
Jaret Burkett
fa8fc32c0a Corrected key saving and loading to better match kohya 2023-09-04 00:22:34 -06:00
Jaret Burkett
22ed539321 Allow special args for schedulers 2023-09-03 20:38:44 -06:00
Jaret Burkett
7cd6945082 Added my annotator/preprocessor and improved network jitter on reference trainer 2023-09-03 16:43:51 -06:00
Jaret Burkett
2a40937b4f reworked samplers. Trying to find what is wrong with diffusers sampling is sdxl 2023-09-03 07:56:09 -06:00
Jaret Burkett
4ca819a05e Fixes for dataloader 2023-08-31 04:54:10 -06:00
Jaret Burkett
addf024630 Fixed issue with omitting square pictures 2023-08-30 15:00:22 -06:00
Jaret Burkett
33267e117c Reworked bucket loader to scale buckets to pixels amounts not just minimum size. Makes the network more consistant 2023-08-30 14:52:12 -06:00
Jaret Burkett
d401348c2e Make data loader resiliant to bad headers in meta 2023-08-29 18:56:06 -06:00
Jaret Burkett
836fee47a6 Fixed some mismatched weights by adjusting tolerance. The mismatch ironically made the models better lol 2023-08-29 15:20:03 -06:00
Jaret Burkett
14ff51ceb4 fixed issues with converting and saving models. Cleaned keys. Improved testing for cycle load saving. 2023-08-29 12:31:19 -06:00
Jaret Burkett
714854ee86 Hude rework to move the batch to a DTO to make it far more modular to the future ui 2023-08-29 10:22:19 -06:00
Jaret Burkett
bd758ff203 Cleanup and small bug fixes 2023-08-29 05:45:49 -06:00
Jaret Burkett
a008d9e63b Fixed issue with loadin models after resume function added. Added additional flush if not training text encoder to clear out vram before grad accum 2023-08-28 17:56:30 -06:00
Jaret Burkett
b79ced3e10 Merge branch 'main' into development 2023-08-28 16:21:51 -06:00
Jaret Burkett
bee0b6a235 Added converters for all stable diffusion models to convert back to ldm format from diffusers. 2023-08-28 16:12:32 -06:00
Jaret Burkett
2ecb5cf024 Merge branch 'main' into development 2023-08-28 14:01:44 -06:00
Jaret Burkett
fab7c2b04a Fixed issue with key mapping from diffusers back to ldm 2023-08-28 14:01:26 -06:00
Jaret Burkett
e866c75638 Built base interfaces for a DTO to handle batch infomation transports for the dataloader 2023-08-28 12:43:31 -06:00
Jaret Burkett
71da78c8af improved normalization for a network with varrying batch network weights 2023-08-28 12:42:57 -06:00
Jaret Burkett
c446f768ea Huge memory optimizations, many big fixes 2023-08-27 17:48:02 -06:00
Jaret Burkett
cc49786ee9 Dataloader bug fixes 2023-08-27 14:36:38 -06:00
Jaret Burkett
9b164a8688 Fixed issue with bucket dataloader corpping in too much. Added normalization capabilities to LoRA modules. Testing effects, but should prevent them from burning and also make them more compatable with stacking many LoRAs 2023-08-27 09:40:01 -06:00
Jaret Burkett
6bd3851058 Fixed issue with prompt token replace adding more than one replacement 2023-08-26 18:52:23 -06:00
Jaret Burkett
fd338e67bb Fixed bug with dataloader not seperating mulitple datasets 2023-08-26 18:07:24 -06:00
Jaret Burkett
8105c05c12 Added bucketting capabilities to dataloader. Finally have full planned capability. noice 2023-08-26 16:36:32 -06:00
Jaret Burkett
2cb27c3f57 Merge branch 'main' of github.com:ostris/ai-toolkit 2023-08-26 08:55:09 -06:00
Jaret Burkett
3367ab6b2c Moved SD batch processing to a shared method and added it for use in slider training. Still testing if it affects quality over sampling 2023-08-26 08:55:00 -06:00
Jaret Burkett
24f46ea7d6 Merge pull request #9 from FoundSol/patch-1
Update README.md
2023-08-25 20:05:51 -06:00
Fundamentum
5bef2985b5 Update README.md 2023-08-25 21:13:24 -03:00
Jaret Burkett
aeaca13d69 Fixed issue with shuffeling permutations 2023-08-23 22:02:00 -06:00
Jaret Burkett
b408f9f3eb Fixed issue with timestep I broke for sliders 2023-08-23 16:15:30 -06:00
Jaret Burkett
7157c316af Added support for training lora, dreambooth, and fine tuning. Still need testing and docs 2023-08-23 15:37:00 -06:00
Jaret Burkett
e2c547f6c2 Fixed typo 2023-08-23 13:33:48 -06:00
Jaret Burkett
7b770bc305 Merge pull request #7 from ostris/textual_inversion
Textual inversion training
2023-08-23 13:31:37 -06:00
Jaret Burkett
f200cf36c5 Added train example to ti 2023-08-23 13:30:29 -06:00
Jaret Burkett
d298240cec Tied in ant tested TI script 2023-08-23 13:26:28 -06:00
Jaret Burkett
2e6c55c720 WIP creating textual inversion training script 2023-08-22 21:02:38 -06:00
Jaret Burkett
36ba08d3fa Added a converter back to ldm from diffusers for sdxl. Can finally get to training it properly 2023-08-21 16:22:01 -06:00
Jaret Burkett
e8667f856f Fix issue with there being an extra . on gene 2023-08-20 15:54:38 -06:00
Jaret Burkett
bef5551ea5 Ultimate slider training built, still needs tuning 2023-08-19 18:54:34 -06:00
Jaret Burkett
b77b9acc0b Added base for ultimate slider. WIP 2023-08-19 15:35:24 -06:00
Jaret Burkett
c6675e2801 Added shuffeling to prompts 2023-08-19 07:57:30 -06:00
Jaret Burkett
90eedb78bf Added multiplier jitter, min_snr, ability to choose sdxl encoders to use, shuffle generator, and other fun 2023-08-19 05:54:22 -06:00
Jaret Burkett
80e2f4a2a4 Merge branch 'main' of github.com:ostris/ai-toolkit 2023-08-18 11:45:00 -06:00
Jaret Burkett
d51c4ca704 Added ability to use two seperate folders for datasets when doing image reference sliders 2023-08-18 11:44:33 -06:00
Jaret Burkett
c7ec132d5d third times a charm 2023-08-16 20:54:00 -06:00
Jaret Burkett
ed9607e8da Update SliderTraining.ipynb
I should be doing this the right way
2023-08-16 20:49:09 -06:00
Jaret Burkett
d44f8ac508 Update SliderTraining.ipynb
fixed code block
2023-08-16 20:47:56 -06:00
Jaret Burkett
8d09eb44ec Fixed an issue with CFG time embeds on SDXL 2023-08-15 18:02:14 -06:00
Jaret Burkett
55a5fcc7d9 Added method to get specific keys from model 2023-08-15 14:51:04 -06:00
Jaret Burkett
e96874241d Added slider colab to readme 2023-08-13 13:55:38 -06:00
Jaret Burkett
e3be1a1758 Added WIP slider training colab 2023-08-13 13:52:38 -06:00
Jaret Burkett
1a92e97c6d Added missing deps 2023-08-13 13:15:04 -06:00
Jaret Burkett
355c80df07 Added ability to use civit ai url ar model name and built a model downloader and cache manager for it 2023-08-13 13:09:51 -06:00
Jaret Burkett
1487d13191 Moved the run job command 2023-08-13 10:25:56 -06:00
Jaret Burkett
383bad958d Added a way to run as a library by passing job dict 2023-08-13 09:54:39 -06:00
Jaret Burkett
196b693cf0 Worked on reference slider script. It is working well currently. Still going to tune it a bit before a writeup though 2023-08-12 17:59:24 -06:00
Jaret Burkett
fd95e7b60c Merge branch 'main' of github.com:ostris/ai-toolkit 2023-08-12 05:59:58 -06:00
Jaret Burkett
379992d89e Various bug fixes and improvements 2023-08-12 05:59:50 -06:00
Jaret Burkett
c7054d714f Finally muted the annoying safety checker notification 2023-08-12 01:38:18 -06:00
Jaret Burkett
67dfd9ced0 Added inbuild plugins and made one for image referenced. WIP 2023-08-10 16:20:38 -06:00
Jaret Burkett
1a7e346b41 Added inbuild plugins and made one for image referenced. WIP 2023-08-10 16:02:44 -06:00
Jaret Burkett
df48f0a843 Moved some of the job config into base process so it will be easier to extend extensions 2023-08-10 12:14:05 -06:00
Jaret Burkett
fbc8a87a05 Reworked the sd rescaler script 2023-08-09 08:57:27 -06:00
Jaret Burkett
bf90740b59 Fixed numerous issues with traing ESRGAN 2023-08-08 20:03:19 -06:00
Jaret Burkett
ff2c9f3d04 Merge branch 'main' of github.com:ostris/ai-toolkit 2023-08-07 18:05:03 -06:00
Jaret Burkett
8bd536df7e Added training for a custom version of ERSGAN arcitecture. Testing training now 2023-08-07 18:04:23 -06:00
Jaret Burkett
64fbd4c92a Fix some windows dependency issues 2023-08-05 20:28:49 -06:00
Jaret Burkett
8c90fa86c6 Complete reqork of how slider training works and optimized it to hell. Can run entire algorythm in 1 batch now with less VRAM consumption than a quarter of it used to take 2023-08-05 18:46:08 -06:00
Jaret Burkett
7e4e660663 Added extensions and an example extension that merges models 2023-08-04 09:37:24 -06:00
Jaret Burkett
b865ac8b24 Various windows bug fixes 2023-08-04 05:51:58 -06:00
Jaret Burkett
66c6f0f6f7 Big refactor of SD runner and added image generator 2023-08-03 14:51:25 -06:00
Jaret Burkett
75ec5d9292 Hotfix to handle latest transformers clip model missing key suddenly 2023-08-02 12:26:18 -06:00
Jaret Burkett
1a25b275c8 Did some work on SD rescaler. Need to run a long test on it eventually. 2023-08-02 07:59:27 -06:00
Jaret Burkett
2bf3e529ce Set gradient checkpointing on unet enabled by default. Help out immensly with sdxl backprop spikes 2023-08-01 15:43:27 -06:00
Jaret Burkett
f53fd08690 Fixed issue with rescaled loras only saving af fp32 2023-08-01 14:08:22 -06:00
Jaret Burkett
8b8d53888d Added Model rescale and prepared a release upgrade 2023-08-01 13:49:54 -06:00
Jaret Burkett
63cacf4362 Merge remote-tracking branch 'origin/main' into WIP 2023-07-31 15:14:28 -06:00
Jaret Burkett
c1b1e800df Updated module names to be compatable with ney koyha unet 2023-07-31 13:16:29 -06:00
Jaret Burkett
7726911562 Allow name to be passed through command line 2023-07-31 11:56:57 -06:00
Jaret Burkett
c01673f1b5 Added random weight adjuster to prevent overfitting 2023-07-29 19:30:14 -06:00
Jaret Burkett
c35b78f0d4 Added random weight adjuster to prevent overfitting 2023-07-29 17:14:14 -06:00
Jaret Burkett
8ba1b11557 Merge branch 'sdxl' into WIP
# Conflicts:
#	jobs/process/BaseSDTrainProcess.py
#	jobs/process/TrainSliderProcess.py
2023-07-29 14:29:18 -06:00
Jaret Burkett
1e50b39442 Work on slider rework 2023-07-28 18:11:10 -06:00
Jaret Burkett
5fc2bb5d9c Information trainer 2023-07-28 08:16:29 -06:00
Jaret Burkett
c7640b0865 WIP diffusers pipeline is weird. Starting to hate sdxl 2023-07-27 17:35:24 -06:00
Jaret Burkett
b2e2e4bf47 Added sd1.5 and 2.1 do the diffusers pipeline flow 2023-07-27 12:34:48 -06:00
Jaret Burkett
596e57a6a6 Pipelines working on SDXL for noise prediction 2023-07-27 11:24:33 -06:00
Jaret Burkett
6ab8b8b0f1 WIP. just need to put it here 2023-07-27 01:46:30 -06:00
186 changed files with 57592 additions and 2583 deletions

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

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

5
.github/ISSUE_TEMPLATE/config.yml vendored Normal file
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.

5
.gitignore vendored
View File

@@ -170,4 +170,7 @@ cython_debug/
!/config/examples
!/config/_PUT_YOUR_CONFIGS_HERE).txt
/output/*
!/output/.gitkeep
!/output/.gitkeep
/extensions/*
!/extensions/example
/temp

6
.gitmodules vendored
View File

@@ -4,3 +4,9 @@
[submodule "repositories/leco"]
path = repositories/leco
url = https://github.com/p1atdev/LECO
[submodule "repositories/batch_annotator"]
path = repositories/batch_annotator
url = https://github.com/ostris/batch-annotator
[submodule "repositories/ipadapter"]
path = repositories/ipadapter
url = https://github.com/tencent-ailab/IP-Adapter.git

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.

313
README.md
View File

@@ -1,16 +1,21 @@
# AI Toolkit by Ostris
## IMPORTANT NOTE - READ THIS
This is an active WIP repo that is not ready for others to use. And definitely not ready for non developers to use.
I am making major breaking changes and pushing straight to master until I have it in a planned state. I have big changes
planned for config files and the general structure. I may change how training works entirely. You are welcome to use
but keep that in mind. If more people start to use it, I will follow better branch checkout standards, but for now
this is my personal active experiment.
This is my research repo. I do a lot of experiments in it and it is possible that I will break things.
If something breaks, checkout an earlier commit. This repo can train a lot of things, and it is
hard to keep up with all of them.
Report bugs as you find them, but not knowing how to train ML models, setup an environment, or use python is not a bug.
I will make all of this more user-friendly eventually
## Support my work
I will make a better readme later.
<a href="https://glif.app" target="_blank">
<img alt="glif.app" src="https://raw.githubusercontent.com/ostris/ai-toolkit/main/assets/glif.svg?v=1" width="256" height="auto">
</a>
My work on this project would not be possible without the amazing support of [Glif](https://glif.app/) and everyone on the
team. If you want to support me, support Glif. [Join the site](https://glif.app/),
[Join us on Discord](https://discord.com/invite/nuR9zZ2nsh), [follow us on Twitter](https://x.com/heyglif)
and come make some cool stuff with us
## Installation
@@ -29,16 +34,211 @@ cd ai-toolkit
git submodule update --init --recursive
python3 -m venv venv
source venv/bin/activate
# or source venv/Scripts/activate on windows
# .\venv\Scripts\activate on windows
# install torch first
pip3 install torch
pip3 install -r requirements.txt
```
Windows:
```bash
git clone https://github.com/ostris/ai-toolkit.git
cd ai-toolkit
git submodule update --init --recursive
python -m venv venv
.\venv\Scripts\activate
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121
pip install -r requirements.txt
```
## FLUX.1 Training
### Tutorial
To get started quickly, check out [@araminta_k](https://x.com/araminta_k) tutorial on [Finetuning Flux Dev on a 3090](https://www.youtube.com/watch?v=HzGW_Kyermg) with 24GB VRAM.
### Requirements
You currently need a GPU with **at least 24GB of VRAM** to train FLUX.1. If you are using it as your GPU to control
your monitors, you probably need to set the flag `low_vram: true` in the config file under `model:`. This will quantize
the model on CPU and should allow it to train with monitors attached. Users have gotten it to work on Windows with WSL,
but there are some reports of a bug when running on windows natively.
I have only tested on linux for now. This is still extremely experimental
and a lot of quantizing and tricks had to happen to get it to fit on 24GB at all.
### FLUX.1-dev
FLUX.1-dev has a non-commercial license. Which means anything you train will inherit the
non-commercial license. It is also a gated model, so you need to accept the license on HF before using it.
Otherwise, this will fail. Here are the required steps to setup a license.
1. Sign into HF and accept the model access here [black-forest-labs/FLUX.1-dev](https://huggingface.co/black-forest-labs/FLUX.1-dev)
2. Make a file named `.env` in the root on this folder
3. [Get a READ key from huggingface](https://huggingface.co/settings/tokens/new?) and add it to the `.env` file like so `HF_TOKEN=your_key_here`
### FLUX.1-schnell
FLUX.1-schnell is Apache 2.0. Anything trained on it can be licensed however you want and it does not require a HF_TOKEN to train.
However, it does require a special adapter to train with it, [ostris/FLUX.1-schnell-training-adapter](https://huggingface.co/ostris/FLUX.1-schnell-training-adapter).
It is also highly experimental. For best overall quality, training on FLUX.1-dev is recommended.
To use it, You just need to add the assistant to the `model` section of your config file like so:
```yaml
model:
name_or_path: "black-forest-labs/FLUX.1-schnell"
assistant_lora_path: "ostris/FLUX.1-schnell-training-adapter"
is_flux: true
quantize: true
```
You also need to adjust your sample steps since schnell does not require as many
```yaml
sample:
guidance_scale: 1 # schnell does not do guidance
sample_steps: 4 # 1 - 4 works well
```
### Training
1. Copy the example config file located at `config/examples/train_lora_flux_24gb.yaml` (`config/examples/train_lora_flux_schnell_24gb.yaml` for schnell) to the `config` folder and rename it to `whatever_you_want.yml`
2. Edit the file following the comments in the file
3. Run the file like so `python run.py config/whatever_you_want.yml`
A folder with the name and the training folder from the config file will be created when you start. It will have all
checkpoints and images in it. You can stop the training at any time using ctrl+c and when you resume, it will pick back up
from the last checkpoint.
IMPORTANT. If you press crtl+c while it is saving, it will likely corrupt that checkpoint. So wait until it is done saving
### Need help?
Please do not open a bug report unless it is a bug in the code. You are welcome to [Join my Discord](https://discord.gg/VXmU2f5WEU)
and ask for help there. However, please refrain from PMing me directly with general question or support. Ask in the discord
and I will answer when I can.
## Training in RunPod
Example RunPod template: **runpod/pytorch:2.2.0-py3.10-cuda12.1.1-devel-ubuntu22.04**
> You need a minimum of 24GB VRAM, pick a GPU by your preference.
#### Example config ($0.5/hr):
- 1x A40 (48 GB VRAM)
- 19 vCPU 100 GB RAM
#### Custom overrides (you need some storage to clone FLUX.1, store datasets, store trained models and samples):
- ~120 GB Disk
- ~120 GB Pod Volume
- Start Jupyter Notebook
### 1. Setup
```
git clone https://github.com/ostris/ai-toolkit.git
cd ai-toolkit
git submodule update --init --recursive
python -m venv venv
source venv/bin/activate
pip install torch
pip install -r requirements.txt
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
```
### 2. Upload your dataset
- Create a new folder in the root, name it `dataset` or whatever you like.
- Drag and drop your .jpg, .jpeg, or .png images and .txt files inside the newly created dataset folder.
### 3. Login into Hugging Face with an Access Token
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
- Run ```huggingface-cli login``` and paste your token.
### 4. Training
- Copy an example config file located at ```config/examples``` to the config folder and rename it to ```whatever_you_want.yml```.
- Edit the config following the comments in the file.
- Change ```folder_path: "/path/to/images/folder"``` to your dataset path like ```folder_path: "/workspace/ai-toolkit/your-dataset"```.
- Run the file: ```python run.py config/whatever_you_want.yml```.
### Screenshot from RunPod
<img width="1728" alt="RunPod Training Screenshot" src="https://github.com/user-attachments/assets/53a1b8ef-92fa-4481-81a7-bde45a14a7b5">
## Training in Modal
### 1. Setup
#### ai-toolkit:
```
git clone https://github.com/ostris/ai-toolkit.git
cd ai-toolkit
git submodule update --init --recursive
python -m venv venv
source venv/bin/activate
pip install torch
pip install -r requirements.txt
pip install --upgrade accelerate transformers diffusers huggingface_hub #Optional, run it if you run into issues
```
#### Modal:
- Run `pip install modal` to install the modal Python package.
- Run `modal setup` to authenticate (if this doesn’t work, try `python -m modal setup`).
#### Hugging Face:
- Get a READ token from [here](https://huggingface.co/settings/tokens) and request access to Flux.1-dev model from [here](https://huggingface.co/black-forest-labs/FLUX.1-dev).
- Run `huggingface-cli login` and paste your token.
### 2. Upload your dataset
- Drag and drop your dataset folder containing the .jpg, .jpeg, or .png images and .txt files in `ai-toolkit`.
### 3. Configs
- Copy an example config file located at ```config/examples/modal``` to the `config` folder and rename it to ```whatever_you_want.yml```.
- Edit the config following the comments in the file, **<ins>be careful and follow the example `/root/ai-toolkit` paths</ins>**.
### 4. Edit run_modal.py
- Set your entire local `ai-toolkit` path at `code_mount = modal.Mount.from_local_dir` like:
```
code_mount = modal.Mount.from_local_dir("/Users/username/ai-toolkit", remote_path="/root/ai-toolkit")
```
- Choose a `GPU` and `Timeout` in `@app.function` _(default is A100 40GB and 2 hour timeout)_.
### 5. Training
- Run the config file in your terminal: `modal run run_modal.py --config-file-list-str=/root/ai-toolkit/config/whatever_you_want.yml`.
- You can monitor your training in your local terminal, or on [modal.com](https://modal.com/).
- Models, samples and optimizer will be stored in `Storage > flux-lora-models`.
### 6. Saving the model
- Check contents of the volume by running `modal volume ls flux-lora-models`.
- Download the content by running `modal volume get flux-lora-models your-model-name`.
- Example: `modal volume get flux-lora-models my_first_flux_lora_v1`.
### Screenshot from Modal
<img width="1728" alt="Modal Traning Screenshot" src="https://github.com/user-attachments/assets/7497eb38-0090-49d6-8ad9-9c8ea7b5388b">
---
## Current Tools
## Dataset Preparation
I have so many hodge podge scripts I am going to be moving over to this that I use in my ML work. But this is what is
here so far.
Datasets generally need to be a folder containing images and associated text files. Currently, the only supported
formats are jpg, jpeg, and png. Webp currently has issues. The text files should be named the same as the images
but with a `.txt` extension. For example `image2.jpg` and `image2.txt`. The text file should contain only the caption.
You can add the word `[trigger]` in the caption file and if you have `trigger_word` in your config, it will be automatically
replaced.
Images are never upscaled but they are downscaled and placed in buckets for batching. **You do not need to crop/resize your images**.
The loader will automatically resize them and can handle varying aspect ratios.
---
## EVERYTHING BELOW THIS LINE IS OUTDATED
It may still work like that, but I have not tested it in a while.
---
### Batch Image Generation
A image generator that can take frompts from a config file or form a txt file and generate them to a
folder. I mainly needed this for an SDXL test I am doing but added some polish to it so it can be used
for generat batch image generation.
It all runs off a config file, which you can find an example of in `config/examples/generate.example.yaml`.
Mere info is in the comments in the example
---
### LoRA (lierla), LoCON (LyCORIS) extractor
@@ -64,9 +264,38 @@ Most people used fixed, which is traditional fixed dimension extraction.
`process` is an array of different processes to run. You can add a few and mix and match. One LoRA, one LyCON, etc.
---
### LoRA Rescale
Change `<lora:my_lora:4.6>` to `<lora:my_lora:1.0>` or whatever you want with the same effect.
A tool for rescaling a LoRA's weights. Should would with LoCON as well, but I have not tested it.
It all runs off a config file, which you can find an example of in `config/examples/mod_lora_scale.yml`.
Just copy that file, into the `config` folder, and rename it to `whatever_you_want.yml`.
Then you can edit the file to your liking. and call it like so:
```bash
python3 run.py config/whatever_you_want.yml
```
You can also put a full path to a config file, if you want to keep it somewhere else.
```bash
python3 run.py "/home/user/whatever_you_want.yml"
```
More notes on how it works are available in the example config file itself. This is useful when making
all LoRAs, as the ideal weight is rarely 1.0, but now you can fix that. For sliders, they can have weird scales form -2 to 2
or even -15 to 15. This will allow you to dile it in so they all have your desired scale
---
### LoRA Slider Trainer
<a target="_blank" href="https://colab.research.google.com/github/ostris/ai-toolkit/blob/main/notebooks/SliderTraining.ipynb">
<img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/>
</a>
This is how I train most of the recent sliders I have on Civitai, you can check them out in my [Civitai profile](https://civitai.com/user/Ostris/models).
It is based off the work by [p1atdev/LECO](https://github.com/p1atdev/LECO) and [rohitgandikota/erasing](https://github.com/rohitgandikota/erasing)
But has been heavily modified to create sliders rather than erasing concepts. I have a lot more plans on this, but it is
@@ -89,6 +318,23 @@ I will post an better tutorial soon.
---
## Extensions!!
You can now make and share custom extensions. That run within this framework and have all the inbuilt tools
available to them. I will probably use this as the primary development method going
forward so I dont keep adding and adding more and more features to this base repo. I will likely migrate a lot
of the existing functionality as well to make everything modular. There is an example extension in the `extensions`
folder that shows how to make a model merger extension. All of the code is heavily documented which is hopefully
enough to get you started. To make an extension, just copy that example and replace all the things you need to.
### Model Merger - Example Extension
It is located in the `extensions` folder. It is a fully finctional model merger that can merge as many models together
as you want. It is a good example of how to make an extension, but is also a pretty useful feature as well since most
mergers can only do one model at a time and this one will take as many as you want to feed it. There is an
example config file in there, just copy that to your `config` folder and rename it to `whatever_you_want.yml`.
and use it like any other config file.
## WIP Tools
@@ -108,14 +354,53 @@ Just went in and out. It is much worse on smaller faces than shown here.
## TODO
- [X] Add proper regs on sliders
- [ ] Add SDXL support (base model only for now)
- [X] Add SDXL support (base model only for now)
- [ ] Add plain erasing
- [ ] Make Textual inversion network trainer (network that spits out TI embeddings)
---
## Change Log
#### 2021-07-30
#### 2023-08-05
- Huge memory rework and slider rework. Slider training is better thant ever with no more
ram spikes. I also made it so all 4 parts of the slider algorythm run in one batch so they share gradient
accumulation. This makes it much faster and more stable.
- Updated the example config to be something more practical and more updated to current methods. It is now
a detail slide and shows how to train one without a subject. 512x512 slider training for 1.5 should work on
6GB gpu now. Will test soon to verify.
#### 2021-10-20
- Windows support bug fixes
- Extensions! Added functionality to make and share custom extensions for training, merging, whatever.
check out the example in the `extensions` folder. Read more about that above.
- Model Merging, provided via the example extension.
#### 2023-08-03
Another big refactor to make SD more modular.
Made batch image generation script
#### 2023-08-01
Major changes and update. New LoRA rescale tool, look above for details. Added better metadata so
Automatic1111 knows what the base model is. Added some experiments and a ton of updates. This thing is still unstable
at the moment, so hopefully there are not breaking changes.
Unfortunately, I am too lazy to write a proper changelog with all the changes.
I added SDXL training to sliders... but.. it does not work properly.
The slider training relies on a model's ability to understand that an unconditional (negative prompt)
means you do not want that concept in the output. SDXL does not understand this for whatever reason,
which makes separating out
concepts within the model hard. I am sure the community will find a way to fix this
over time, but for now, it is not
going to work properly. And if any of you are thinking "Could we maybe fix it by adding 1 or 2 more text
encoders to the model as well as a few more entirely separate diffusion networks?" No. God no. It just needs a little
training without every experimental new paper added to it. The KISS principal.
#### 2023-07-30
Added "anchors" to the slider trainer. This allows you to set a prompt that will be used as a
regularizer. You can set the network multiplier to force spread consistency at high weights

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

View File

@@ -0,0 +1,8 @@
#!/usr/bin/env bash
echo "Docker builds from the repo, not this dir. Make sure changes are pushed to the repo."
# wait 2 seconds
sleep 2
docker build --build-arg CACHEBUST=$(date +%s) -t aitoolkit:latest -f docker/Dockerfile .
docker tag aitoolkit:latest ostris/aitoolkit:latest
docker push ostris/aitoolkit:latest

View File

@@ -0,0 +1,60 @@
---
job: generate # tells the runner what to do
config:
name: "generate" # this is not really used anywhere currently but required by runner
process:
# process 1
- type: to_folder # process images to a folder
output_folder: "output/gen"
device: cuda:0 # cpu, cuda:0, etc
generate:
# these are your defaults you can override most of them with flags
sampler: "ddpm" # ignored for now, will add later though ddpm is used regardless for now
width: 1024
height: 1024
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
seed: -1 # -1 is random
guidance_scale: 7
sample_steps: 20
ext: ".png" # .png, .jpg, .jpeg, .webp
# here ate the flags you can use for prompts. Always start with
# your prompt first then add these flags after. You can use as many
# like
# photo of a baseball --n painting, ugly --w 1024 --h 1024 --seed 42 --cfg 7 --steps 20
# we will try to support all sd-scripts flags where we can
# FROM SD-SCRIPTS
# --n Treat everything until the next option as a negative prompt.
# --w Specify the width of the generated image.
# --h Specify the height of the generated image.
# --d Specify the seed for the generated image.
# --l Specify the CFG scale for the generated image.
# --s Specify the number of steps during generation.
# OURS and some QOL additions
# --p2 Prompt for the second text encoder (SDXL only)
# --n2 Negative prompt for the second text encoder (SDXL only)
# --gr Specify the guidance rescale for the generated image (SDXL only)
# --seed Specify the seed for the generated image same as --d
# --cfg Specify the CFG scale for the generated image same as --l
# --steps Specify the number of steps during generation same as --s
prompt_file: false # if true a txt file will be created next to images with prompt strings used
# prompts can also be a path to a text file with one prompt per line
# prompts: "/path/to/prompts.txt"
prompts:
- "photo of batman"
- "photo of superman"
- "photo of spiderman"
- "photo of a superhero --n batman superman spiderman"
model:
# huggingface name, relative prom project path, or absolute path to .safetensors or .ckpt
# name_or_path: "runwayml/stable-diffusion-v1-5"
name_or_path: "/mnt/Models/stable-diffusion/models/stable-diffusion/Ostris/Ostris_Real_v1.safetensors"
is_v2: false # for v2 models
is_v_pred: false # for v-prediction models (most v2 models)
is_xl: false # for SDXL models
dtype: bf16

View File

@@ -0,0 +1,48 @@
---
job: mod
config:
name: name_of_your_model_v1
process:
- type: rescale_lora
# path to your current lora model
input_path: "/path/to/lora/lora.safetensors"
# output path for your new lora model, can be the same as input_path to replace
output_path: "/path/to/lora/output_lora_v1.safetensors"
# replaces meta with the meta below (plus minimum meta fields)
# if false, we will leave the meta alone except for updating hashes (sd-script hashes)
replace_meta: true
# how to adjust, we can scale the up_down weights or the alpha
# up_down is the default and probably the best, they will both net the same outputs
# would only affect rare NaN cases and maybe merging with old merge tools
scale_target: 'up_down'
# precision to save, fp16 is the default and standard
save_dtype: fp16
# current_weight is the ideal weight you use as a multiplier when using the lora
# IE in automatic1111 <lora:my_lora:6.0> the 6.0 is the current_weight
# you can do negatives here too if you want to flip the lora
current_weight: 6.0
# target_weight is the ideal weight you use as a multiplier when using the lora
# instead of the one above. IE in automatic1111 instead of using <lora:my_lora:6.0>
# we want to use <lora:my_lora:1.0> so 1.0 is the target_weight
target_weight: 1.0
# base model for the lora
# this is just used to add meta so automatic111 knows which model it is for
# assume v1.5 if these are not set
is_xl: false
is_v2: false
meta:
# this is only used if you set replace_meta to true above
name: "[name]" # [name] gets replaced with the name above
description: A short description of your lora
trigger_words:
- put
- trigger
- words
- here
version: '0.1'
creator:
name: Your Name
email: your@email.com
website: https://yourwebsite.com
any: All meta data above is arbitrary, it can be whatever you want.

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

@@ -7,7 +7,7 @@ job: train
config:
# the name will be used to create a folder in the output folder
# it will also replace any [name] token in the rest of this config
name: pet_slider_v1
name: detail_slider_v1
# folder will be created with name above in folder below
# it can be relative to the project root or absolute
training_folder: "output/LoRA"
@@ -23,9 +23,8 @@ config:
# network type lierla is traditional LoRA that works everywhere, only linear layers
type: "lierla"
# rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good
rank: 8
alpha: 1.0 # just leave it
linear: 8
linear_alpha: 4 # Do about half of rank
# training config
train:
# this is also used in sampling. Stick with ddpm unless you know what you are doing
@@ -33,14 +32,17 @@ config:
# how many steps to train. More is not always better. I rarely go over 1000
steps: 500
# I have had good results with 4e-4 to 1e-4 at 500 steps
lr: 1e-4
lr: 2e-4
# enables gradient checkpoint, saves vram, leave it on
gradient_checkpointing: true
# train the unet. I recommend leaving this true
train_unet: true
# train the text encoder. I don't recommend this unless you have a special use case
# for sliders we are adjusting representation of the concept (unet),
# not the description of it (text encoder)
train_text_encoder: false
# same as from sd-scripts, not fully tested but should speed up training
min_snr_gamma: 5.0
# just leave unless you know what you are doing
# also supports "dadaptation" but set lr to 1 if you use that,
# but it learns too fast and I don't recommend it
@@ -51,14 +53,17 @@ config:
# while training. Just leave it
max_denoising_steps: 40
# works great at 1. I do 1 even with my 4090.
# higher may not work right with newer single batch stacking code anyway
batch_size: 1
# bf16 works best if your GPU supports it (modern)
dtype: bf16 # fp32, bf16, fp16
# if you have it, use it. It is faster and better
xformers: true
# torch 2.0 doesnt need xformers anymore, only use if you have lower version
# xformers: true
# I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX
# although, the way we train sliders is comparative, so it probably won't work anyway
noise_offset: 0.0
# noise_offset: 0.0357 # SDXL was trained with offset of 0.0357. So use that when training on SDXL
# the model to train the LoRA network on
model:
@@ -66,11 +71,17 @@ config:
name_or_path: "runwayml/stable-diffusion-v1-5"
is_v2: false # for v2 models
is_v_pred: false # for v-prediction models (most v2 models)
# has some issues with the dual text encoder and the way we train sliders
# it works bit weights need to probably be higher to see it.
is_xl: false # for SDXL models
# saving config
save:
dtype: float16 # precision to save. I recommend float16
save_every: 50 # save every this many steps
# this will remove step counts more than this number
# allows you to save more often in case of a crash without filling up your drive
max_step_saves_to_keep: 2
# sampling config
sample:
@@ -88,21 +99,22 @@ config:
# --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive
# slide are good tests. will inherit sample.network_multiplier if not set
# --n [string] # negative prompt, will inherit sample.neg if not set
# Only 75 tokens allowed currently
prompts: # our example is an animal slider, neg: dog, pos: cat
- "a golden retriever --m -5"
- "a golden retriever --m -3"
- "a golden retriever --m 3"
- "a golden retriever --m 5"
- "calico cat --m -5"
- "calico cat --m -3"
- "calico cat --m 3"
- "calico cat --m 5"
- "an elephant --m -5"
- "an elephant --m -3"
- "an elephant --m 3"
- "an elephant --m 5"
# I like to do a wide positive and negative spread so I can see a good range and stop
# early if the network is braking down
prompts:
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5"
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3"
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3"
- "a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5"
- "a golden retriever sitting on a leather couch, --m -5"
- "a golden retriever sitting on a leather couch --m -3"
- "a golden retriever sitting on a leather couch --m 3"
- "a golden retriever sitting on a leather couch --m 5"
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5"
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3"
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3"
- "a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5"
# negative prompt used on all prompts above as default if they don't have one
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome"
# seed for sampling. 42 is the answer for everything
@@ -131,11 +143,16 @@ config:
# resolutions to train on. [ width, height ]. This is less important for sliders
# as we are not teaching the model anything it doesn't already know
# but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1
# and [ 1024, 1024 ] for sd_xl
# you can do as many as you want here
resolutions:
- [ 512, 512 ]
# - [ 512, 768 ]
# - [ 768, 768 ]
# slider training uses 4 combined steps for a single round. This will do it in one gradient
# step. It is highly optimized and shouldn't take anymore vram than doing without it,
# since we break down batches for gradient accumulation now. so just leave it on.
batch_full_slide: true
# These are the concepts to train on. You can do as many as you want here,
# but they can conflict outweigh each other. Other than experimenting, I recommend
# just doing one for good results
@@ -146,7 +163,9 @@ config:
# a keyword necessarily but what the model understands the concept to represent.
# "person" will affect men, women, children, etc but will not affect cats, dogs, etc
# it is the models base general understanding of the concept and everything it represents
- target_class: "animal"
# you can leave it blank to affect everything. In this example, we are adjusting
# detail, so we will leave it blank to affect everything
- target_class: ""
# positive is the prompt for the positive side of the slider.
# It is the concept that will be excited and amplified in the model when we slide the slider
# to the positive side and forgotten / inverted when we slide
@@ -154,33 +173,48 @@ config:
# the prompt. You want it to be the extreme of what you want to train on. For example,
# if you want to train on fat people, you would use "an extremely fat, morbidly obese person"
# as the prompt. Not just "fat person"
positive: "cat"
# max 75 tokens for now
positive: "high detail, 8k, intricate, detailed, high resolution, high res, high quality"
# negative is the prompt for the negative side of the slider and works the same as positive
# it does not necessarily work the same as a negative prompt when generating images
negative: "dog"
# these need to be polar opposites.
# max 76 tokens for now
negative: "blurry, boring, fuzzy, low detail, low resolution, low res, low quality"
# the loss for this target is multiplied by this number.
# if you are doing more than one target it may be good to set less important ones
# to a lower number like 0.1 so they dont outweigh the primary target
# to a lower number like 0.1 so they don't outweigh the primary target
weight: 1.0
# shuffle the prompts split by the comma. We will run every combination randomly
# this will make the LoRA more robust. You probably want this on unless prompt order
# is important for some reason
shuffle: true
# anchors are prompts that wer try to hold on to while training the slider
# you want these to generate an image very similar to the target_class
# without directly overlapping it. For example, if you are training on a person smiling,
# you would use "a person with a face mask" as an anchor. It is a person, the image is the same
# regardless if they are smiling or not
anchors:
# only positive prompt for now
- prompt: "a woman"
neg_prompt: "animal"
# the multiplier applied to the LoRA when this is run.
# higher will give it more weight but also help keep the lora from collapsing
multiplier: 8.0
- prompt: "a man"
neg_prompt: "animal"
multiplier: 8.0
- prompt: "a person"
neg_prompt: "animal"
multiplier: 8.0
# anchors are prompts that we will try to hold on to while training the slider
# these are NOT necessary and can prevent the slider from converging if not done right
# leave them off if you are having issues, but they can help lock the network
# on certain concepts to help prevent catastrophic forgetting
# you want these to generate an image that is not your target_class, but close to it
# is fine as long as it does not directly overlap it.
# For example, if you are training on a person smiling,
# you could use "a person with a face mask" as an anchor. It is a person, the image is the same
# regardless if they are smiling or not, however, the closer the concept is to the target_class
# the less the multiplier needs to be. Keep multipliers less than 1.0 for anchors usually
# for close concepts, you want to be closer to 0.1 or 0.2
# these will slow down training. I am leaving them off for the demo
# anchors:
# - prompt: "a woman"
# neg_prompt: "animal"
# # the multiplier applied to the LoRA when this is run.
# # higher will give it more weight but also help keep the lora from collapsing
# multiplier: 1.0
# - prompt: "a man"
# neg_prompt: "animal"
# multiplier: 1.0
# - prompt: "a person"
# neg_prompt: "animal"
# multiplier: 1.0
# You can put any information you want here, and it will be saved in the model.
# The below is an example, but you can put your grocery list in it if you want.

21
docker/Dockerfile Normal file
View File

@@ -0,0 +1,21 @@
FROM runpod/base:0.6.2-cuda12.1.0
LABEL authors="jaret"
# Install dependencies
RUN apt-get update
WORKDIR /app
ARG CACHEBUST=1
RUN git clone https://github.com/ostris/ai-toolkit.git && \
cd ai-toolkit && \
git submodule update --init --recursive
WORKDIR /app/ai-toolkit
RUN ln -s /usr/bin/python3 /usr/bin/python
RUN python -m pip install -r requirements.txt
RUN apt-get install -y tmux nvtop htop
WORKDIR /
CMD ["/start.sh"]

View File

@@ -0,0 +1,129 @@
import torch
import gc
from collections import OrderedDict
from typing import TYPE_CHECKING
from jobs.process import BaseExtensionProcess
from toolkit.config_modules import ModelConfig
from toolkit.stable_diffusion_model import StableDiffusion
from toolkit.train_tools import get_torch_dtype
from tqdm import tqdm
# Type check imports. Prevents circular imports
if TYPE_CHECKING:
from jobs import ExtensionJob
# extend standard config classes to add weight
class ModelInputConfig(ModelConfig):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.weight = kwargs.get('weight', 1.0)
# overwrite default dtype unless user specifies otherwise
# float 32 will give up better precision on the merging functions
self.dtype: str = kwargs.get('dtype', 'float32')
def flush():
torch.cuda.empty_cache()
gc.collect()
# this is our main class process
class ExampleMergeModels(BaseExtensionProcess):
def __init__(
self,
process_id: int,
job: 'ExtensionJob',
config: OrderedDict
):
super().__init__(process_id, job, config)
# this is the setup process, do not do process intensive stuff here, just variable setup and
# checking requirements. This is called before the run() function
# no loading models or anything like that, it is just for setting up the process
# all of your process intensive stuff should be done in the run() function
# config will have everything from the process item in the config file
# convince methods exist on BaseProcess to get config values
# if required is set to true and the value is not found it will throw an error
# you can pass a default value to get_conf() as well if it was not in the config file
# as well as a type to cast the value to
self.save_path = self.get_conf('save_path', required=True)
self.save_dtype = self.get_conf('save_dtype', default='float16', as_type=get_torch_dtype)
self.device = self.get_conf('device', default='cpu', as_type=torch.device)
# build models to merge list
models_to_merge = self.get_conf('models_to_merge', required=True, as_type=list)
# build list of ModelInputConfig objects. I find it is a good idea to make a class for each config
# this way you can add methods to it and it is easier to read and code. There are a lot of
# inbuilt config classes located in toolkit.config_modules as well
self.models_to_merge = [ModelInputConfig(**model) for model in models_to_merge]
# setup is complete. Don't load anything else here, just setup variables and stuff
# this is the entire run process be sure to call super().run() first
def run(self):
# always call first
super().run()
print(f"Running process: {self.__class__.__name__}")
# let's adjust our weights first to normalize them so the total is 1.0
total_weight = sum([model.weight for model in self.models_to_merge])
weight_adjust = 1.0 / total_weight
for model in self.models_to_merge:
model.weight *= weight_adjust
output_model: StableDiffusion = None
# let's do the merge, it is a good idea to use tqdm to show progress
for model_config in tqdm(self.models_to_merge, desc="Merging models"):
# setup model class with our helper class
sd_model = StableDiffusion(
device=self.device,
model_config=model_config,
dtype="float32"
)
# load the model
sd_model.load_model()
# adjust the weight of the text encoder
if isinstance(sd_model.text_encoder, list):
# sdxl model
for text_encoder in sd_model.text_encoder:
for key, value in text_encoder.state_dict().items():
value *= model_config.weight
else:
# normal model
for key, value in sd_model.text_encoder.state_dict().items():
value *= model_config.weight
# adjust the weights of the unet
for key, value in sd_model.unet.state_dict().items():
value *= model_config.weight
if output_model is None:
# use this one as the base
output_model = sd_model
else:
# merge the models
# text encoder
if isinstance(output_model.text_encoder, list):
# sdxl model
for i, text_encoder in enumerate(output_model.text_encoder):
for key, value in text_encoder.state_dict().items():
value += sd_model.text_encoder[i].state_dict()[key]
else:
# normal model
for key, value in output_model.text_encoder.state_dict().items():
value += sd_model.text_encoder.state_dict()[key]
# unet
for key, value in output_model.unet.state_dict().items():
value += sd_model.unet.state_dict()[key]
# remove the model to free memory
del sd_model
flush()
# merge loop is done, let's save the model
print(f"Saving merged model to {self.save_path}")
output_model.save(self.save_path, meta=self.meta, save_dtype=self.save_dtype)
print(f"Saved merged model to {self.save_path}")
# do cleanup here
del output_model
flush()

View File

@@ -0,0 +1,25 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# We make a subclass of Extension
class ExampleMergeExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "example_merge_extension"
# name is the name of the extension for printing
name = "Example Merge Extension"
# 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 .ExampleMergeModels import ExampleMergeModels
return ExampleMergeModels
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
ExampleMergeExtension
]

View File

@@ -0,0 +1,48 @@
---
# Always include at least one example config file to show how to use your extension.
# use plenty of comments so users know how to use it and what everything does
# all extensions will use this job name
job: extension
config:
name: 'my_awesome_merge'
process:
# Put your example processes here. This will be passed
# to your extension process in the config argument.
# the type MUST match your extension uid
- type: "example_merge_extension"
# save path for the merged model
save_path: "output/merge/[name].safetensors"
# save type
dtype: fp16
# device to run it on
device: cuda:0
# input models can only be SD1.x and SD2.x models for this example (currently)
models_to_merge:
# weights are relative, total weights will be normalized
# for example. If you have 2 models with weight 1.0, they will
# both be weighted 0.5. If you have 1 model with weight 1.0 and
# another with weight 2.0, the first will be weighted 1/3 and the
# second will be weighted 2/3
- name_or_path: "input/model1.safetensors"
weight: 1.0
- name_or_path: "input/model2.safetensors"
weight: 1.0
- name_or_path: "input/model3.safetensors"
weight: 0.3
- name_or_path: "input/model4.safetensors"
weight: 1.0
# you can put any information you want here, and it will be saved in the model
# the below is an example. I recommend doing trigger words at a minimum
# in the metadata. The software will include this plus some other information
meta:
name: "[name]" # [name] gets replaced with the name above
description: A short description of your model
version: '0.1'
creator:
name: Your Name
email: your@email.com
website: https://yourwebsite.com
any: All meta data above is arbitrary, it can be whatever you want.

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

@@ -0,0 +1,102 @@
import os
from collections import OrderedDict
from toolkit.config_modules import ModelConfig, GenerateImageConfig, SampleConfig, LoRMConfig
from toolkit.lorm import ExtractMode, convert_diffusers_unet_to_lorm
from toolkit.sd_device_states_presets import get_train_sd_device_state_preset
from toolkit.stable_diffusion_model import StableDiffusion
import gc
import torch
from jobs.process import BaseExtensionProcess
from toolkit.train_tools import get_torch_dtype
def flush():
torch.cuda.empty_cache()
gc.collect()
class PureLoraGenerator(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.device = self.get_conf('device', 'cuda')
self.device_torch = torch.device(self.device)
self.model_config = ModelConfig(**self.get_conf('model', required=True))
self.generate_config = SampleConfig(**self.get_conf('sample', required=True))
self.dtype = self.get_conf('dtype', 'float16')
self.torch_dtype = get_torch_dtype(self.dtype)
lorm_config = self.get_conf('lorm', None)
self.lorm_config = LoRMConfig(**lorm_config) if lorm_config is not None else None
self.device_state_preset = get_train_sd_device_state_preset(
device=torch.device(self.device),
)
self.progress_bar = None
self.sd = StableDiffusion(
device=self.device,
model_config=self.model_config,
dtype=self.dtype,
)
def run(self):
super().run()
print("Loading model...")
with torch.no_grad():
self.sd.load_model()
self.sd.unet.eval()
self.sd.unet.to(self.device_torch)
if isinstance(self.sd.text_encoder, list):
for te in self.sd.text_encoder:
te.eval()
te.to(self.device_torch)
else:
self.sd.text_encoder.eval()
self.sd.to(self.device_torch)
print(f"Converting to LoRM UNet")
# replace the unet with LoRMUnet
convert_diffusers_unet_to_lorm(
self.sd.unet,
config=self.lorm_config,
)
sample_folder = os.path.join(self.output_folder)
gen_img_config_list = []
sample_config = self.generate_config
start_seed = sample_config.seed
current_seed = start_seed
for i in range(len(sample_config.prompts)):
if sample_config.walk_seed:
current_seed = start_seed + i
filename = f"[time]_[count].{self.generate_config.ext}"
output_path = os.path.join(sample_folder, filename)
prompt = sample_config.prompts[i]
extra_args = {}
gen_img_config_list.append(GenerateImageConfig(
prompt=prompt, # it will autoparse the prompt
width=sample_config.width,
height=sample_config.height,
negative_prompt=sample_config.neg,
seed=current_seed,
guidance_scale=sample_config.guidance_scale,
guidance_rescale=sample_config.guidance_rescale,
num_inference_steps=sample_config.sample_steps,
network_multiplier=sample_config.network_multiplier,
output_path=output_path,
output_ext=sample_config.ext,
adapter_conditioning_scale=sample_config.adapter_conditioning_scale,
**extra_args
))
# send to be generated
self.sd.generate_images(gen_img_config_list, sampler=sample_config.sampler)
print("Done generating images")
# cleanup
del self.sd
gc.collect()
torch.cuda.empty_cache()

View File

@@ -0,0 +1,212 @@
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 torch.utils.data import DataLoader
from diffusers import StableDiffusionXLAdapterPipeline, StableDiffusionAdapterPipeline
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
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.t2i_adapter_path = kwargs.get('t2i_adapter_path', None)
self.guidance_scale = kwargs.get('guidance_scale', 7)
self.sample_steps = kwargs.get('sample_steps', 20)
self.prompt_2 = kwargs.get('prompt_2', None)
self.neg_2 = kwargs.get('neg_2', None)
self.prompts = kwargs.get('prompts', None)
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
self.ext = kwargs.get('ext', 'png')
self.adapter_conditioning_scale = kwargs.get('adapter_conditioning_scale', 1.0)
if kwargs.get('shuffle', False):
# shuffle the prompts
random.shuffle(self.prompts)
class ReferenceGenerator(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.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 run(self):
super().run()
print("Loading model...")
self.sd.load_model()
device = torch.device(self.device)
if self.generate_config.t2i_adapter_path is not None:
self.adapter = T2IAdapter.from_pretrained(
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)
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)
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
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]
output_path = os.path.join(self.output_folder, img_filename)
output_caption_path = os.path.join(self.output_folder, img_filename_no_ext + '.txt')
output_depth_path = os.path.join(self.output_folder, img_filename_no_ext + '.depth.png')
caption = batch.get_caption_list()[0]
img: torch.Tensor = batch.tensor.clone()
# 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)
width, height = image.size
min_res = min(width, height)
if self.generate_config.walk_seed:
seed = seed + 1
if self.generate_config.seed == -1:
# random
seed = random.randint(0, 1000000)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
# generate depth map
image = midas_depth(
image,
detect_resolution=min_res, # do 512 ?
image_resolution=min_res
)
# image.save(output_depth_path)
gen_images = pipe(
prompt=caption,
negative_prompt=self.generate_config.neg,
image=image,
num_inference_steps=self.generate_config.sample_steps,
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
with open(output_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

@@ -0,0 +1,59 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# This is for generic training (LoRA, Dreambooth, FineTuning)
class AdvancedReferenceGeneratorExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "reference_generator"
# name is the name of the extension for printing
name = "Reference Generator"
# 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 .ReferenceGenerator import ReferenceGenerator
return ReferenceGenerator
# This is for generic training (LoRA, Dreambooth, FineTuning)
class PureLoraGenerator(Extension):
# uid must be unique, it is how the extension is identified
uid = "pure_lora_generator"
# name is the name of the extension for printing
name = "Pure LoRA Generator"
# 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 .PureLoraGenerator import PureLoraGenerator
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, Img2ImgGeneratorExtension
]

View File

@@ -0,0 +1,91 @@
---
job: extension
config:
name: test_v1
process:
- type: 'textual_inversion_trainer'
training_folder: "out/TI"
device: cuda:0
# for tensorboard logging
log_dir: "out/.tensorboard"
embedding:
trigger: "your_trigger_here"
tokens: 12
init_words: "man with short brown hair"
save_format: "safetensors" # 'safetensors' or 'pt'
save:
dtype: float16 # precision to save
save_every: 100 # save every this many steps
max_step_saves_to_keep: 5 # only affects step counts
datasets:
- folder_path: "/path/to/dataset"
caption_ext: "txt"
default_caption: "[trigger]"
buckets: true
resolution: 512
train:
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
steps: 3000
weight_jitter: 0.0
lr: 5e-5
train_unet: false
gradient_checkpointing: true
train_text_encoder: false
optimizer: "adamw"
# optimizer: "prodigy"
optimizer_params:
weight_decay: 1e-2
lr_scheduler: "constant"
max_denoising_steps: 1000
batch_size: 4
dtype: bf16
xformers: true
min_snr_gamma: 5.0
# skip_first_sample: true
noise_offset: 0.0 # not needed for this
model:
# objective reality v2
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
is_v2: false # for v2 models
is_xl: false # for SDXL models
is_v_pred: false # for v-prediction models (most v2 models)
sample:
sampler: "ddpm" # must match train.noise_scheduler
sample_every: 100 # sample every this many steps
width: 512
height: 512
prompts:
- "photo of [trigger] laughing"
- "photo of [trigger] smiling"
- "[trigger] close up"
- "dark scene [trigger] frozen"
- "[trigger] nighttime"
- "a painting of [trigger]"
- "a drawing of [trigger]"
- "a cartoon of [trigger]"
- "[trigger] pixar style"
- "[trigger] costume"
neg: ""
seed: 42
walk_seed: false
guidance_scale: 7
sample_steps: 20
network_multiplier: 1.0
logging:
log_every: 10 # log every this many steps
use_wandb: false # not supported yet
verbose: false
# You can put any information you want here, and it will be saved in the model.
# The below is an example, but you can put your grocery list in it if you want.
# It is saved in the model so be aware of that. The software will include this
# plus some other information for you automatically
meta:
# [name] gets replaced with the name above
name: "[name]"
# version: '1.0'
# creator:
# name: Your Name
# email: your@gmail.com
# website: https://your.website

View File

@@ -0,0 +1,151 @@
import random
from collections import OrderedDict
from torch.utils.data import DataLoader
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds
from toolkit.stable_diffusion_model import StableDiffusion, BlankNetwork
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
import gc
import torch
from jobs.process import BaseSDTrainProcess
def flush():
torch.cuda.empty_cache()
gc.collect()
class ConceptReplacementConfig:
def __init__(self, **kwargs):
self.concept: str = kwargs.get('concept', '')
self.replacement: str = kwargs.get('replacement', '')
class ConceptReplacer(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super().__init__(process_id, job, config, **kwargs)
replacement_list = self.config.get('replacements', [])
self.replacement_list = [ConceptReplacementConfig(**x) for x in replacement_list]
def before_model_load(self):
pass
def hook_before_train_loop(self):
self.sd.vae.eval()
self.sd.vae.to(self.device_torch)
# textual inversion
if self.embedding is not None:
# set text encoder to train. Not sure if this is necessary but diffusers example did it
self.sd.text_encoder.train()
def hook_train_loop(self, batch):
with torch.no_grad():
dtype = get_torch_dtype(self.train_config.dtype)
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
network_weight_list = batch.get_network_weight_list()
# have a blank network so we can wrap it in a context and set multipliers without checking every time
if self.network is not None:
network = self.network
else:
network = BlankNetwork()
batch_replacement_list = []
# get a random replacement for each prompt
for prompt in conditioned_prompts:
replacement = random.choice(self.replacement_list)
batch_replacement_list.append(replacement)
# build out prompts
concept_prompts = []
replacement_prompts = []
for idx, replacement in enumerate(batch_replacement_list):
prompt = conditioned_prompts[idx]
# insert shuffled concept at beginning and end of prompt
shuffled_concept = [x.strip() for x in replacement.concept.split(',')]
random.shuffle(shuffled_concept)
shuffled_concept = ', '.join(shuffled_concept)
concept_prompts.append(f"{shuffled_concept}, {prompt}, {shuffled_concept}")
# insert replacement at beginning and end of prompt
shuffled_replacement = [x.strip() for x in replacement.replacement.split(',')]
random.shuffle(shuffled_replacement)
shuffled_replacement = ', '.join(shuffled_replacement)
replacement_prompts.append(f"{shuffled_replacement}, {prompt}, {shuffled_replacement}")
# predict the replacement without network
conditional_embeds = self.sd.encode_prompt(replacement_prompts).to(self.device_torch, dtype=dtype)
replacement_pred = self.sd.predict_noise(
latents=noisy_latents.to(self.device_torch, dtype=dtype),
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
timestep=timesteps,
guidance_scale=1.0,
)
del conditional_embeds
replacement_pred = replacement_pred.detach()
self.optimizer.zero_grad()
flush()
# text encoding
grad_on_text_encoder = False
if self.train_config.train_text_encoder:
grad_on_text_encoder = True
if self.embedding:
grad_on_text_encoder = True
# set the weights
network.multiplier = network_weight_list
# activate network if it exits
with network:
with torch.set_grad_enabled(grad_on_text_encoder):
# embed the prompts
conditional_embeds = self.sd.encode_prompt(concept_prompts).to(self.device_torch, dtype=dtype)
if not grad_on_text_encoder:
# detach the embeddings
conditional_embeds = conditional_embeds.detach()
self.optimizer.zero_grad()
flush()
noise_pred = self.sd.predict_noise(
latents=noisy_latents.to(self.device_torch, dtype=dtype),
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
timestep=timesteps,
guidance_scale=1.0,
)
loss = torch.nn.functional.mse_loss(noise_pred.float(), replacement_pred.float(), reduction="none")
loss = loss.mean([1, 2, 3])
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
# add min_snr_gamma
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler, self.train_config.min_snr_gamma)
loss = loss.mean()
# back propagate loss to free ram
loss.backward()
flush()
# apply gradients
self.optimizer.step()
self.optimizer.zero_grad()
self.lr_scheduler.step()
if self.embedding is not None:
# Let's make sure we don't update any embedding weights besides the newly added token
self.embedding.restore_embeddings()
loss_dict = OrderedDict(
{'loss': loss.item()}
)
# reset network multiplier
network.multiplier = 1.0
return loss_dict

View File

@@ -0,0 +1,26 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# This is for generic training (LoRA, Dreambooth, FineTuning)
class ConceptReplacerExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "concept_replacer"
# name is the name of the extension for printing
name = "Concept Replacer"
# 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 .ConceptReplacer import ConceptReplacer
return ConceptReplacer
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
ConceptReplacerExtension,
]

View File

@@ -0,0 +1,91 @@
---
job: extension
config:
name: test_v1
process:
- type: 'textual_inversion_trainer'
training_folder: "out/TI"
device: cuda:0
# for tensorboard logging
log_dir: "out/.tensorboard"
embedding:
trigger: "your_trigger_here"
tokens: 12
init_words: "man with short brown hair"
save_format: "safetensors" # 'safetensors' or 'pt'
save:
dtype: float16 # precision to save
save_every: 100 # save every this many steps
max_step_saves_to_keep: 5 # only affects step counts
datasets:
- folder_path: "/path/to/dataset"
caption_ext: "txt"
default_caption: "[trigger]"
buckets: true
resolution: 512
train:
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
steps: 3000
weight_jitter: 0.0
lr: 5e-5
train_unet: false
gradient_checkpointing: true
train_text_encoder: false
optimizer: "adamw"
# optimizer: "prodigy"
optimizer_params:
weight_decay: 1e-2
lr_scheduler: "constant"
max_denoising_steps: 1000
batch_size: 4
dtype: bf16
xformers: true
min_snr_gamma: 5.0
# skip_first_sample: true
noise_offset: 0.0 # not needed for this
model:
# objective reality v2
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
is_v2: false # for v2 models
is_xl: false # for SDXL models
is_v_pred: false # for v-prediction models (most v2 models)
sample:
sampler: "ddpm" # must match train.noise_scheduler
sample_every: 100 # sample every this many steps
width: 512
height: 512
prompts:
- "photo of [trigger] laughing"
- "photo of [trigger] smiling"
- "[trigger] close up"
- "dark scene [trigger] frozen"
- "[trigger] nighttime"
- "a painting of [trigger]"
- "a drawing of [trigger]"
- "a cartoon of [trigger]"
- "[trigger] pixar style"
- "[trigger] costume"
neg: ""
seed: 42
walk_seed: false
guidance_scale: 7
sample_steps: 20
network_multiplier: 1.0
logging:
log_every: 10 # log every this many steps
use_wandb: false # not supported yet
verbose: false
# You can put any information you want here, and it will be saved in the model.
# The below is an example, but you can put your grocery list in it if you want.
# It is saved in the model so be aware of that. The software will include this
# plus some other information for you automatically
meta:
# [name] gets replaced with the name above
name: "[name]"
# version: '1.0'
# creator:
# name: Your Name
# email: your@gmail.com
# website: https://your.website

View File

@@ -0,0 +1,20 @@
from collections import OrderedDict
import gc
import torch
from jobs.process import BaseExtensionProcess
def flush():
torch.cuda.empty_cache()
gc.collect()
class DatasetTools(BaseExtensionProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
def run(self):
super().run()
raise NotImplementedError("This extension is not yet implemented")

View File

@@ -0,0 +1,196 @@
import copy
import json
import os
from collections import OrderedDict
import gc
import traceback
import torch
from PIL import Image, ImageOps
from tqdm import tqdm
from .tools.dataset_tools_config_modules import RAW_DIR, TRAIN_DIR, Step, ImgInfo
from .tools.fuyu_utils import FuyuImageProcessor
from .tools.image_tools import load_image, ImageProcessor, resize_to_max
from .tools.llava_utils import LLaVAImageProcessor
from .tools.caption import default_long_prompt, default_short_prompt, default_replacements
from jobs.process import BaseExtensionProcess
from .tools.sync_tools import get_img_paths
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
def flush():
torch.cuda.empty_cache()
gc.collect()
VERSION = 2
class SuperTagger(BaseExtensionProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
parent_dir = config.get('parent_dir', None)
self.dataset_paths: list[str] = config.get('dataset_paths', [])
self.device = config.get('device', 'cuda')
self.steps: list[Step] = config.get('steps', [])
self.caption_method = config.get('caption_method', 'llava:default')
self.caption_prompt = config.get('caption_prompt', default_long_prompt)
self.caption_short_prompt = config.get('caption_short_prompt', default_short_prompt)
self.force_reprocess_img = config.get('force_reprocess_img', False)
self.caption_replacements = config.get('caption_replacements', default_replacements)
self.caption_short_replacements = config.get('caption_short_replacements', default_replacements)
self.master_dataset_dict = OrderedDict()
self.dataset_master_config_file = config.get('dataset_master_config_file', None)
if parent_dir is not None and len(self.dataset_paths) == 0:
# find all folders in the patent_dataset_path
self.dataset_paths = [
os.path.join(parent_dir, folder)
for folder in os.listdir(parent_dir)
if os.path.isdir(os.path.join(parent_dir, folder))
]
else:
# make sure they exist
for dataset_path in self.dataset_paths:
if not os.path.exists(dataset_path):
raise ValueError(f"Dataset path does not exist: {dataset_path}")
print(f"Found {len(self.dataset_paths)} dataset paths")
self.image_processor: ImageProcessor = self.get_image_processor()
def get_image_processor(self):
if self.caption_method.startswith('llava'):
return LLaVAImageProcessor(device=self.device)
elif self.caption_method.startswith('fuyu'):
return FuyuImageProcessor(device=self.device)
else:
raise ValueError(f"Unknown caption method: {self.caption_method}")
def process_image(self, img_path: str):
root_img_dir = os.path.dirname(os.path.dirname(img_path))
filename = os.path.basename(img_path)
filename_no_ext = os.path.splitext(filename)[0]
train_dir = os.path.join(root_img_dir, TRAIN_DIR)
train_img_path = os.path.join(train_dir, filename)
json_path = os.path.join(train_dir, f"{filename_no_ext}.json")
# check if json exists, if it does load it as image info
if os.path.exists(json_path):
with open(json_path, 'r') as f:
img_info = ImgInfo(**json.load(f))
else:
img_info = ImgInfo()
# always send steps first in case other processes need them
img_info.add_steps(copy.deepcopy(self.steps))
img_info.set_version(VERSION)
img_info.set_caption_method(self.caption_method)
image: Image = None
caption_image: Image = None
did_update_image = False
# trigger reprocess of steps
if self.force_reprocess_img:
img_info.trigger_image_reprocess()
# set the image as updated if it does not exist on disk
if not os.path.exists(train_img_path):
did_update_image = True
image = load_image(img_path)
if img_info.force_image_process:
did_update_image = True
image = load_image(img_path)
# go through the needed steps
for step in copy.deepcopy(img_info.state.steps_to_complete):
if step == 'caption':
# load image
if image is None:
image = load_image(img_path)
if caption_image is None:
caption_image = resize_to_max(image, 1024, 1024)
if not self.image_processor.is_loaded:
print('Loading Model. Takes a while, especially the first time')
self.image_processor.load_model()
img_info.caption = self.image_processor.generate_caption(
image=caption_image,
prompt=self.caption_prompt,
replacements=self.caption_replacements
)
img_info.mark_step_complete(step)
elif step == 'caption_short':
# load image
if image is None:
image = load_image(img_path)
if caption_image is None:
caption_image = resize_to_max(image, 1024, 1024)
if not self.image_processor.is_loaded:
print('Loading Model. Takes a while, especially the first time')
self.image_processor.load_model()
img_info.caption_short = self.image_processor.generate_caption(
image=caption_image,
prompt=self.caption_short_prompt,
replacements=self.caption_short_replacements
)
img_info.mark_step_complete(step)
elif step == 'contrast_stretch':
# load image
if image is None:
image = load_image(img_path)
image = ImageOps.autocontrast(image, cutoff=(0.1, 0), preserve_tone=True)
did_update_image = True
img_info.mark_step_complete(step)
else:
raise ValueError(f"Unknown step: {step}")
os.makedirs(os.path.dirname(train_img_path), exist_ok=True)
if did_update_image:
image.save(train_img_path)
if img_info.is_dirty:
with open(json_path, 'w') as f:
json.dump(img_info.to_dict(), f, indent=4)
if self.dataset_master_config_file:
# add to master dict
self.master_dataset_dict[train_img_path] = img_info.to_dict()
def run(self):
super().run()
imgs_to_process = []
# find all images
for dataset_path in self.dataset_paths:
raw_dir = os.path.join(dataset_path, RAW_DIR)
raw_image_paths = get_img_paths(raw_dir)
for raw_image_path in raw_image_paths:
imgs_to_process.append(raw_image_path)
if len(imgs_to_process) == 0:
print(f"No images to process")
else:
print(f"Found {len(imgs_to_process)} to process")
for img_path in tqdm(imgs_to_process, desc="Processing images"):
try:
self.process_image(img_path)
except Exception:
# print full stack trace
print(traceback.format_exc())
continue
# self.process_image(img_path)
if self.dataset_master_config_file is not None:
# save it as json
with open(self.dataset_master_config_file, 'w') as f:
json.dump(self.master_dataset_dict, f, indent=4)
del self.image_processor
flush()

View File

@@ -0,0 +1,131 @@
import os
import shutil
from collections import OrderedDict
import gc
from typing import List
import torch
from tqdm import tqdm
from .tools.dataset_tools_config_modules import DatasetSyncCollectionConfig, RAW_DIR, NEW_DIR
from .tools.sync_tools import get_unsplash_images, get_pexels_images, get_local_image_file_names, download_image, \
get_img_paths
from jobs.process import BaseExtensionProcess
def flush():
torch.cuda.empty_cache()
gc.collect()
class SyncFromCollection(BaseExtensionProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.min_width = config.get('min_width', 1024)
self.min_height = config.get('min_height', 1024)
# add our min_width and min_height to each dataset config if they don't exist
for dataset_config in config.get('dataset_sync', []):
if 'min_width' not in dataset_config:
dataset_config['min_width'] = self.min_width
if 'min_height' not in dataset_config:
dataset_config['min_height'] = self.min_height
self.dataset_configs: List[DatasetSyncCollectionConfig] = [
DatasetSyncCollectionConfig(**dataset_config)
for dataset_config in config.get('dataset_sync', [])
]
print(f"Found {len(self.dataset_configs)} dataset configs")
def move_new_images(self, root_dir: str):
raw_dir = os.path.join(root_dir, RAW_DIR)
new_dir = os.path.join(root_dir, NEW_DIR)
new_images = get_img_paths(new_dir)
for img_path in new_images:
# move to raw
new_path = os.path.join(raw_dir, os.path.basename(img_path))
shutil.move(img_path, new_path)
# remove new dir
shutil.rmtree(new_dir)
def sync_dataset(self, config: DatasetSyncCollectionConfig):
if config.host == 'unsplash':
get_images = get_unsplash_images
elif config.host == 'pexels':
get_images = get_pexels_images
else:
raise ValueError(f"Unknown host: {config.host}")
results = {
'num_downloaded': 0,
'num_skipped': 0,
'bad': 0,
'total': 0,
}
photos = get_images(config)
raw_dir = os.path.join(config.directory, RAW_DIR)
new_dir = os.path.join(config.directory, NEW_DIR)
raw_images = get_local_image_file_names(raw_dir)
new_images = get_local_image_file_names(new_dir)
for photo in tqdm(photos, desc=f"{config.host}-{config.collection_id}"):
try:
if photo.filename not in raw_images and photo.filename not in new_images:
download_image(photo, new_dir, min_width=self.min_width, min_height=self.min_height)
results['num_downloaded'] += 1
else:
results['num_skipped'] += 1
except Exception as e:
print(f" - BAD({photo.id}): {e}")
results['bad'] += 1
continue
results['total'] += 1
return results
def print_results(self, results):
print(
f" - new:{results['num_downloaded']}, old:{results['num_skipped']}, bad:{results['bad']} total:{results['total']}")
def run(self):
super().run()
print(f"Syncing {len(self.dataset_configs)} datasets")
all_results = None
failed_datasets = []
for dataset_config in tqdm(self.dataset_configs, desc="Syncing datasets", leave=True):
try:
results = self.sync_dataset(dataset_config)
if all_results is None:
all_results = {**results}
else:
for key, value in results.items():
all_results[key] += value
self.print_results(results)
except Exception as e:
print(f" - FAILED: {e}")
if 'response' in e.__dict__:
error = f"{e.response.status_code}: {e.response.text}"
print(f" - {error}")
failed_datasets.append({'dataset': dataset_config, 'error': error})
else:
failed_datasets.append({'dataset': dataset_config, 'error': str(e)})
continue
print("Moving new images to raw")
for dataset_config in self.dataset_configs:
self.move_new_images(dataset_config.directory)
print("Done syncing datasets")
self.print_results(all_results)
if len(failed_datasets) > 0:
print(f"Failed to sync {len(failed_datasets)} datasets")
for failed in failed_datasets:
print(f" - {failed['dataset'].host}-{failed['dataset'].collection_id}")
print(f" - ERR: {failed['error']}")

View File

@@ -0,0 +1,43 @@
from toolkit.extension import Extension
class DatasetToolsExtension(Extension):
uid = "dataset_tools"
# name is the name of the extension for printing
name = "Dataset Tools"
# 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 .DatasetTools import DatasetTools
return DatasetTools
class SyncFromCollectionExtension(Extension):
uid = "sync_from_collection"
name = "Sync from Collection"
@classmethod
def get_process(cls):
# import your process class here so it is only loaded when needed and return it
from .SyncFromCollection import SyncFromCollection
return SyncFromCollection
class SuperTaggerExtension(Extension):
uid = "super_tagger"
name = "Super Tagger"
@classmethod
def get_process(cls):
# import your process class here so it is only loaded when needed and return it
from .SuperTagger import SuperTagger
return SuperTagger
AI_TOOLKIT_EXTENSIONS = [
SyncFromCollectionExtension, DatasetToolsExtension, SuperTaggerExtension
]

View File

@@ -0,0 +1,53 @@
caption_manipulation_steps = ['caption', 'caption_short']
default_long_prompt = 'caption this image. describe every single thing in the image in detail. Do not include any unnecessary words in your description for the sake of good grammar. I want many short statements that serve the single purpose of giving the most thorough description if items as possible in the smallest, comma separated way possible. be sure to describe people\'s moods, clothing, the environment, lighting, colors, and everything.'
default_short_prompt = 'caption this image in less than ten words'
default_replacements = [
("the image features", ""),
("the image shows", ""),
("the image depicts", ""),
("the image is", ""),
("in this image", ""),
("in the image", ""),
]
def clean_caption(cap, replacements=None):
if replacements is None:
replacements = default_replacements
# remove any newlines
cap = cap.replace("\n", ", ")
cap = cap.replace("\r", ", ")
cap = cap.replace(".", ",")
cap = cap.replace("\"", "")
# remove unicode characters
cap = cap.encode('ascii', 'ignore').decode('ascii')
# make lowercase
cap = cap.lower()
# remove any extra spaces
cap = " ".join(cap.split())
for replacement in replacements:
if replacement[0].startswith('*'):
# we are removing all text if it starts with this and the rest matches
search_text = replacement[0][1:]
if cap.startswith(search_text):
cap = ""
else:
cap = cap.replace(replacement[0].lower(), replacement[1].lower())
cap_list = cap.split(",")
# trim whitespace
cap_list = [c.strip() for c in cap_list]
# remove empty strings
cap_list = [c for c in cap_list if c != ""]
# remove duplicates
cap_list = list(dict.fromkeys(cap_list))
# join back together
cap = ", ".join(cap_list)
return cap

View File

@@ -0,0 +1,187 @@
import json
from typing import Literal, Type, TYPE_CHECKING
Host: Type = Literal['unsplash', 'pexels']
RAW_DIR = "raw"
NEW_DIR = "_tmp"
TRAIN_DIR = "train"
DEPTH_DIR = "depth"
from .image_tools import Step, img_manipulation_steps
from .caption import caption_manipulation_steps
class DatasetSyncCollectionConfig:
def __init__(self, **kwargs):
self.host: Host = kwargs.get('host', None)
self.collection_id: str = kwargs.get('collection_id', None)
self.directory: str = kwargs.get('directory', None)
self.api_key: str = kwargs.get('api_key', None)
self.min_width: int = kwargs.get('min_width', 1024)
self.min_height: int = kwargs.get('min_height', 1024)
if self.host is None:
raise ValueError("host is required")
if self.collection_id is None:
raise ValueError("collection_id is required")
if self.directory is None:
raise ValueError("directory is required")
if self.api_key is None:
raise ValueError(f"api_key is required: {self.host}:{self.collection_id}")
class ImageState:
def __init__(self, **kwargs):
self.steps_complete: list[Step] = kwargs.get('steps_complete', [])
self.steps_to_complete: list[Step] = kwargs.get('steps_to_complete', [])
def to_dict(self):
return {
'steps_complete': self.steps_complete
}
class Rect:
def __init__(self, **kwargs):
self.x = kwargs.get('x', 0)
self.y = kwargs.get('y', 0)
self.width = kwargs.get('width', 0)
self.height = kwargs.get('height', 0)
def to_dict(self):
return {
'x': self.x,
'y': self.y,
'width': self.width,
'height': self.height
}
class ImgInfo:
def __init__(self, **kwargs):
self.version: int = kwargs.get('version', None)
self.caption: str = kwargs.get('caption', None)
self.caption_short: str = kwargs.get('caption_short', None)
self.poi = [Rect(**poi) for poi in kwargs.get('poi', [])]
self.state = ImageState(**kwargs.get('state', {}))
self.caption_method = kwargs.get('caption_method', None)
self.other_captions = kwargs.get('other_captions', {})
self._upgrade_state()
self.force_image_process: bool = False
self._requested_steps: list[Step] = []
self.is_dirty: bool = False
def _upgrade_state(self):
# upgrades older states
if self.caption is not None and 'caption' not in self.state.steps_complete:
self.mark_step_complete('caption')
self.is_dirty = True
if self.caption_short is not None and 'caption_short' not in self.state.steps_complete:
self.mark_step_complete('caption_short')
self.is_dirty = True
if self.caption_method is None and self.caption is not None:
# added caption method in version 2. Was all llava before that
self.caption_method = 'llava:default'
self.is_dirty = True
def to_dict(self):
return {
'version': self.version,
'caption_method': self.caption_method,
'caption': self.caption,
'caption_short': self.caption_short,
'poi': [poi.to_dict() for poi in self.poi],
'state': self.state.to_dict(),
'other_captions': self.other_captions
}
def mark_step_complete(self, step: Step):
if step not in self.state.steps_complete:
self.state.steps_complete.append(step)
if step in self.state.steps_to_complete:
self.state.steps_to_complete.remove(step)
self.is_dirty = True
def add_step(self, step: Step):
if step not in self.state.steps_to_complete and step not in self.state.steps_complete:
self.state.steps_to_complete.append(step)
def trigger_image_reprocess(self):
if self._requested_steps is None:
raise Exception("Must call add_steps before trigger_image_reprocess")
steps = self._requested_steps
# remove all image manipulationf from steps_to_complete
for step in img_manipulation_steps:
if step in self.state.steps_to_complete:
self.state.steps_to_complete.remove(step)
if step in self.state.steps_complete:
self.state.steps_complete.remove(step)
self.force_image_process = True
self.is_dirty = True
# we want to keep the order passed in process file
for step in steps:
if step in img_manipulation_steps:
self.add_step(step)
def add_steps(self, steps: list[Step]):
self._requested_steps = [step for step in steps]
for stage in steps:
self.add_step(stage)
# update steps if we have any img processes not complete, we have to reprocess them all
# if any steps_to_complete are in img_manipulation_steps
is_manipulating_image = any([step in img_manipulation_steps for step in self.state.steps_to_complete])
order_has_changed = False
if not is_manipulating_image:
# check to see if order has changed. No need to if already redoing it. Will detect if ones are removed
target_img_manipulation_order = [step for step in steps if step in img_manipulation_steps]
current_img_manipulation_order = [step for step in self.state.steps_complete if
step in img_manipulation_steps]
if target_img_manipulation_order != current_img_manipulation_order:
order_has_changed = True
if is_manipulating_image or order_has_changed:
self.trigger_image_reprocess()
def set_caption_method(self, method: str):
if self._requested_steps is None:
raise Exception("Must call add_steps before set_caption_method")
if self.caption_method != method:
self.is_dirty = True
# move previous caption method to other_captions
if self.caption_method is not None and self.caption is not None or self.caption_short is not None:
self.other_captions[self.caption_method] = {
'caption': self.caption,
'caption_short': self.caption_short,
}
self.caption_method = method
self.caption = None
self.caption_short = None
# see if we have a caption from the new method
if method in self.other_captions:
self.caption = self.other_captions[method].get('caption', None)
self.caption_short = self.other_captions[method].get('caption_short', None)
else:
self.trigger_new_caption()
def trigger_new_caption(self):
self.caption = None
self.caption_short = None
self.is_dirty = True
# check to see if we have any steps in the complete list and move them to the to_complete list
for step in self.state.steps_complete:
if step in caption_manipulation_steps:
self.state.steps_complete.remove(step)
self.state.steps_to_complete.append(step)
def to_json(self):
return json.dumps(self.to_dict())
def set_version(self, version: int):
if self.version != version:
self.is_dirty = True
self.version = version

View File

@@ -0,0 +1,66 @@
from transformers import CLIPImageProcessor, BitsAndBytesConfig, AutoTokenizer
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
import torch
from PIL import Image
class FuyuImageProcessor:
def __init__(self, device='cuda'):
from transformers import FuyuProcessor, FuyuForCausalLM
self.device = device
self.model: FuyuForCausalLM = None
self.processor: FuyuProcessor = None
self.dtype = torch.bfloat16
self.tokenizer: AutoTokenizer
self.is_loaded = False
def load_model(self):
from transformers import FuyuProcessor, FuyuForCausalLM
model_path = "adept/fuyu-8b"
kwargs = {"device_map": self.device}
kwargs['load_in_4bit'] = True
kwargs['quantization_config'] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=self.dtype,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type='nf4'
)
self.processor = FuyuProcessor.from_pretrained(model_path)
self.model = FuyuForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
self.is_loaded = True
self.tokenizer = AutoTokenizer.from_pretrained(model_path)
self.model = FuyuForCausalLM.from_pretrained(model_path, torch_dtype=self.dtype, **kwargs)
self.processor = FuyuProcessor(image_processor=FuyuImageProcessor(), tokenizer=self.tokenizer)
def generate_caption(
self, image: Image,
prompt: str = default_long_prompt,
replacements=default_replacements,
max_new_tokens=512
):
# prepare inputs for the model
# text_prompt = f"{prompt}\n"
# image = image.convert('RGB')
model_inputs = self.processor(text=prompt, images=[image])
model_inputs = {k: v.to(dtype=self.dtype if torch.is_floating_point(v) else v.dtype, device=self.device) for k, v in
model_inputs.items()}
generation_output = self.model.generate(**model_inputs, max_new_tokens=max_new_tokens)
prompt_len = model_inputs["input_ids"].shape[-1]
output = self.tokenizer.decode(generation_output[0][prompt_len:], skip_special_tokens=True)
output = clean_caption(output, replacements=replacements)
return output
# inputs = self.processor(text=text_prompt, images=image, return_tensors="pt")
# for k, v in inputs.items():
# inputs[k] = v.to(self.device)
# # autoregressively generate text
# generation_output = self.model.generate(**inputs, max_new_tokens=max_new_tokens)
# generation_text = self.processor.batch_decode(generation_output[:, -max_new_tokens:], skip_special_tokens=True)
# output = generation_text[0]
#
# return clean_caption(output, replacements=replacements)

View File

@@ -0,0 +1,49 @@
from typing import Literal, Type, TYPE_CHECKING, Union
import cv2
import numpy as np
from PIL import Image, ImageOps
Step: Type = Literal['caption', 'caption_short', 'create_mask', 'contrast_stretch']
img_manipulation_steps = ['contrast_stretch']
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
if TYPE_CHECKING:
from .llava_utils import LLaVAImageProcessor
from .fuyu_utils import FuyuImageProcessor
ImageProcessor = Union['LLaVAImageProcessor', 'FuyuImageProcessor']
def pil_to_cv2(image):
"""Convert a PIL image to a cv2 image."""
return cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
def cv2_to_pil(image):
"""Convert a cv2 image to a PIL image."""
return Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
def load_image(img_path: str):
image = Image.open(img_path).convert('RGB')
try:
# transpose with exif data
image = ImageOps.exif_transpose(image)
except Exception as e:
pass
return image
def resize_to_max(image, max_width=1024, max_height=1024):
width, height = image.size
if width <= max_width and height <= max_height:
return image
scale = min(max_width / width, max_height / height)
width = int(width * scale)
height = int(height * scale)
return image.resize((width, height), Image.LANCZOS)

View File

@@ -0,0 +1,85 @@
from .caption import default_long_prompt, default_short_prompt, default_replacements, clean_caption
import torch
from PIL import Image, ImageOps
from transformers import AutoTokenizer, BitsAndBytesConfig, CLIPImageProcessor
img_ext = ['.jpg', '.jpeg', '.png', '.webp']
class LLaVAImageProcessor:
def __init__(self, device='cuda'):
try:
from llava.model import LlavaLlamaForCausalLM
except ImportError:
# print("You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
print(
"You need to manually install llava -> pip install --no-deps git+https://github.com/haotian-liu/LLaVA.git")
raise
self.device = device
self.model: LlavaLlamaForCausalLM = None
self.tokenizer: AutoTokenizer = None
self.image_processor: CLIPImageProcessor = None
self.is_loaded = False
def load_model(self):
from llava.model import LlavaLlamaForCausalLM
model_path = "4bit/llava-v1.5-13b-3GB"
# kwargs = {"device_map": "auto"}
kwargs = {"device_map": self.device}
kwargs['load_in_4bit'] = True
kwargs['quantization_config'] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type='nf4'
)
self.model = LlavaLlamaForCausalLM.from_pretrained(model_path, low_cpu_mem_usage=True, **kwargs)
self.tokenizer = AutoTokenizer.from_pretrained(model_path, use_fast=False)
vision_tower = self.model.get_vision_tower()
if not vision_tower.is_loaded:
vision_tower.load_model()
vision_tower.to(device=self.device)
self.image_processor = vision_tower.image_processor
self.is_loaded = True
def generate_caption(
self, image:
Image, prompt: str = default_long_prompt,
replacements=default_replacements,
max_new_tokens=512
):
from llava.conversation import conv_templates, SeparatorStyle
from llava.utils import disable_torch_init
from llava.constants import IMAGE_TOKEN_INDEX, DEFAULT_IMAGE_TOKEN, DEFAULT_IM_START_TOKEN, DEFAULT_IM_END_TOKEN
from llava.mm_utils import tokenizer_image_token, KeywordsStoppingCriteria
# question = "how many dogs are in the picture?"
disable_torch_init()
conv_mode = "llava_v0"
conv = conv_templates[conv_mode].copy()
roles = conv.roles
image_tensor = self.image_processor.preprocess([image], return_tensors='pt')['pixel_values'].half().cuda()
inp = f"{roles[0]}: {prompt}"
inp = DEFAULT_IM_START_TOKEN + DEFAULT_IMAGE_TOKEN + DEFAULT_IM_END_TOKEN + '\n' + inp
conv.append_message(conv.roles[0], inp)
conv.append_message(conv.roles[1], None)
raw_prompt = conv.get_prompt()
input_ids = tokenizer_image_token(raw_prompt, self.tokenizer, IMAGE_TOKEN_INDEX,
return_tensors='pt').unsqueeze(0).cuda()
stop_str = conv.sep if conv.sep_style != SeparatorStyle.TWO else conv.sep2
keywords = [stop_str]
stopping_criteria = KeywordsStoppingCriteria(keywords, self.tokenizer, input_ids)
with torch.inference_mode():
output_ids = self.model.generate(
input_ids, images=image_tensor, do_sample=True, temperature=0.1,
max_new_tokens=max_new_tokens, use_cache=True, stopping_criteria=[stopping_criteria],
top_p=0.8
)
outputs = self.tokenizer.decode(output_ids[0, input_ids.shape[1]:]).strip()
conv.messages[-1][-1] = outputs
output = outputs.rsplit('</s>', 1)[0]
return clean_caption(output, replacements=replacements)

View File

@@ -0,0 +1,279 @@
import os
import requests
import tqdm
from typing import List, Optional, TYPE_CHECKING
def img_root_path(img_id: str):
return os.path.dirname(os.path.dirname(img_id))
if TYPE_CHECKING:
from .dataset_tools_config_modules import DatasetSyncCollectionConfig
img_exts = ['.jpg', '.jpeg', '.webp', '.png']
class Photo:
def __init__(
self,
id,
host,
width,
height,
url,
filename
):
self.id = str(id)
self.host = host
self.width = width
self.height = height
self.url = url
self.filename = filename
def get_desired_size(img_width: int, img_height: int, min_width: int, min_height: int):
if img_width > img_height:
scale = min_height / img_height
else:
scale = min_width / img_width
new_width = int(img_width * scale)
new_height = int(img_height * scale)
return new_width, new_height
def get_pexels_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
all_images = []
next_page = f"https://api.pexels.com/v1/collections/{config.collection_id}?page=1&per_page=80&type=photos"
while True:
response = requests.get(next_page, headers={
"Authorization": f"{config.api_key}"
})
response.raise_for_status()
data = response.json()
all_images.extend(data['media'])
if 'next_page' in data and data['next_page']:
next_page = data['next_page']
else:
break
photos = []
for image in all_images:
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
url = f"{image['src']['original']}?auto=compress&cs=tinysrgb&h={new_height}&w={new_width}"
filename = os.path.basename(image['src']['original'])
photos.append(Photo(
id=image['id'],
host="pexels",
width=image['width'],
height=image['height'],
url=url,
filename=filename
))
return photos
def get_unsplash_images(config: 'DatasetSyncCollectionConfig') -> List[Photo]:
headers = {
# "Authorization": f"Client-ID {UNSPLASH_ACCESS_KEY}"
"Authorization": f"Client-ID {config.api_key}"
}
# headers['Authorization'] = f"Bearer {token}"
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page=1&per_page=30"
response = requests.get(url, headers=headers)
response.raise_for_status()
res_headers = response.headers
# parse the link header to get the next page
# 'Link': '<https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=82>; rel="last", <https://api.unsplash.com/collections/mIPWwLdfct8/photos?page=2>; rel="next"'
has_next_page = False
if 'Link' in res_headers:
has_next_page = True
link_header = res_headers['Link']
link_header = link_header.split(',')
link_header = [link.strip() for link in link_header]
link_header = [link.split(';') for link in link_header]
link_header = [[link[0].strip('<>'), link[1].strip().strip('"')] for link in link_header]
link_header = {link[1]: link[0] for link in link_header}
# get page number from last url
last_page = link_header['rel="last']
last_page = last_page.split('?')[1]
last_page = last_page.split('&')
last_page = [param.split('=') for param in last_page]
last_page = {param[0]: param[1] for param in last_page}
last_page = int(last_page['page'])
all_images = response.json()
if has_next_page:
# assume we start on page 1, so we don't need to get it again
for page in tqdm.tqdm(range(2, last_page + 1)):
url = f"https://api.unsplash.com/collections/{config.collection_id}/photos?page={page}&per_page=30"
response = requests.get(url, headers=headers)
response.raise_for_status()
all_images.extend(response.json())
photos = []
for image in all_images:
new_width, new_height = get_desired_size(image['width'], image['height'], config.min_width, config.min_height)
url = f"{image['urls']['raw']}&w={new_width}"
filename = f"{image['id']}.jpg"
photos.append(Photo(
id=image['id'],
host="unsplash",
width=image['width'],
height=image['height'],
url=url,
filename=filename
))
return photos
def get_img_paths(dir_path: str):
os.makedirs(dir_path, exist_ok=True)
local_files = os.listdir(dir_path)
# remove non image files
local_files = [file for file in local_files if os.path.splitext(file)[1].lower() in img_exts]
# make full path
local_files = [os.path.join(dir_path, file) for file in local_files]
return local_files
def get_local_image_ids(dir_path: str):
os.makedirs(dir_path, exist_ok=True)
local_files = get_img_paths(dir_path)
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
return set([os.path.basename(file).split('.')[0] for file in local_files])
def get_local_image_file_names(dir_path: str):
os.makedirs(dir_path, exist_ok=True)
local_files = get_img_paths(dir_path)
# assuming local files are named after Unsplash IDs, e.g., 'abc123.jpg'
return set([os.path.basename(file) for file in local_files])
def download_image(photo: Photo, dir_path: str, min_width: int = 1024, min_height: int = 1024):
img_width = photo.width
img_height = photo.height
if img_width < min_width or img_height < min_height:
raise ValueError(f"Skipping {photo.id} because it is too small: {img_width}x{img_height}")
img_response = requests.get(photo.url)
img_response.raise_for_status()
os.makedirs(dir_path, exist_ok=True)
filename = os.path.join(dir_path, photo.filename)
with open(filename, 'wb') as file:
file.write(img_response.content)
def update_caption(img_path: str):
# if the caption is a txt file, convert it to a json file
filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]
# see if it exists
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json")):
# todo add poi and what not
return # we have a json file
caption = ""
# see if txt file exists
if os.path.exists(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt")):
# read it
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"), 'r') as file:
caption = file.read()
# write json file
with open(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.json"), 'w') as file:
file.write(f'{{"caption": "{caption}"}}')
# delete txt file
os.remove(os.path.join(os.path.dirname(img_path), f"{filename_no_ext}.txt"))
# def equalize_img(img_path: str):
# input_path = img_path
# output_path = os.path.join(img_root_path(img_path), COLOR_CORRECTED_DIR, os.path.basename(img_path))
# os.makedirs(os.path.dirname(output_path), exist_ok=True)
# process_img(
# img_path=input_path,
# output_path=output_path,
# equalize=True,
# max_size=2056,
# white_balance=False,
# gamma_correction=False,
# strength=0.6,
# )
# def annotate_depth(img_path: str):
# # make fake args
# args = argparse.Namespace()
# args.annotator = "midas"
# args.res = 1024
#
# img = cv2.imread(img_path)
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
#
# output = annotate(img, args)
#
# output = output.astype('uint8')
# output = cv2.cvtColor(output, cv2.COLOR_RGB2BGR)
#
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
# output_path = os.path.join(img_root_path(img_path), DEPTH_DIR, os.path.basename(img_path))
#
# cv2.imwrite(output_path, output)
# def invert_depth(img_path: str):
# img = cv2.imread(img_path)
# img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# # invert the colors
# img = cv2.bitwise_not(img)
#
# os.makedirs(os.path.dirname(img_path), exist_ok=True)
# output_path = os.path.join(img_root_path(img_path), INVERTED_DEPTH_DIR, os.path.basename(img_path))
# cv2.imwrite(output_path, img)
#
# # update our list of raw images
# raw_images = get_img_paths(raw_dir)
#
# # update raw captions
# for image_id in tqdm.tqdm(raw_images, desc="Updating raw captions"):
# update_caption(image_id)
#
# # equalize images
# for img_path in tqdm.tqdm(raw_images, desc="Equalizing images"):
# if img_path not in eq_images:
# equalize_img(img_path)
#
# # update our list of eq images
# eq_images = get_img_paths(eq_dir)
# # update eq captions
# for image_id in tqdm.tqdm(eq_images, desc="Updating eq captions"):
# update_caption(image_id)
#
# # annotate depth
# depth_dir = os.path.join(root_dir, DEPTH_DIR)
# depth_images = get_img_paths(depth_dir)
# for img_path in tqdm.tqdm(eq_images, desc="Annotating depth"):
# if img_path not in depth_images:
# annotate_depth(img_path)
#
# depth_images = get_img_paths(depth_dir)
#
# # invert depth
# inv_depth_dir = os.path.join(root_dir, INVERTED_DEPTH_DIR)
# inv_depth_images = get_img_paths(inv_depth_dir)
# for img_path in tqdm.tqdm(depth_images, desc="Inverting depth"):
# if img_path not in inv_depth_images:
# invert_depth(img_path)

View File

@@ -0,0 +1,235 @@
import copy
import random
from collections import OrderedDict
import os
from contextlib import nullcontext
from typing import Optional, Union, List
from torch.utils.data import ConcatDataset, DataLoader
from toolkit.config_modules import ReferenceDatasetConfig
from toolkit.data_loader import PairedImageDataset
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
import gc
from toolkit import train_tools
import torch
from jobs.process import BaseSDTrainProcess
import random
from toolkit.basic import value_map
def flush():
torch.cuda.empty_cache()
gc.collect()
class ReferenceSliderConfig:
def __init__(self, **kwargs):
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
self.weight_jitter: float = kwargs.get('weight_jitter', 0.0)
self.datasets: List[ReferenceDatasetConfig] = [ReferenceDatasetConfig(**d) for d in kwargs.get('datasets', [])]
class ImageReferenceSliderTrainerProcess(BaseSDTrainProcess):
sd: StableDiffusion
data_loader: DataLoader = None
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super().__init__(process_id, job, config, **kwargs)
self.prompt_txt_list = None
self.step_num = 0
self.start_step = 0
self.device = self.get_conf('device', self.job.device)
self.device_torch = torch.device(self.device)
self.slider_config = ReferenceSliderConfig(**self.get_conf('slider', {}))
def load_datasets(self):
if self.data_loader is None:
print(f"Loading datasets")
datasets = []
for dataset in self.slider_config.datasets:
print(f" - Dataset: {dataset.pair_folder}")
config = {
'path': dataset.pair_folder,
'size': dataset.size,
'default_prompt': dataset.target_class,
'network_weight': dataset.network_weight,
'pos_weight': dataset.pos_weight,
'neg_weight': dataset.neg_weight,
'pos_folder': dataset.pos_folder,
'neg_folder': dataset.neg_folder,
}
image_dataset = PairedImageDataset(config)
datasets.append(image_dataset)
concatenated_dataset = ConcatDataset(datasets)
self.data_loader = DataLoader(
concatenated_dataset,
batch_size=self.train_config.batch_size,
shuffle=True,
num_workers=2
)
def before_model_load(self):
pass
def hook_before_train_loop(self):
self.sd.vae.eval()
self.sd.vae.to(self.device_torch)
self.load_datasets()
pass
def hook_train_loop(self, batch):
with torch.no_grad():
imgs, prompts, network_weights = batch
network_pos_weight, network_neg_weight = network_weights
if isinstance(network_pos_weight, torch.Tensor):
network_pos_weight = network_pos_weight.item()
if isinstance(network_neg_weight, torch.Tensor):
network_neg_weight = network_neg_weight.item()
# get an array of random floats between -weight_jitter and weight_jitter
loss_jitter_multiplier = 1.0
weight_jitter = self.slider_config.weight_jitter
if weight_jitter > 0.0:
jitter_list = random.uniform(-weight_jitter, weight_jitter)
orig_network_pos_weight = network_pos_weight
network_pos_weight += jitter_list
network_neg_weight += (jitter_list * -1.0)
# penalize the loss for its distance from network_pos_weight
# a jitter_list of abs(3.0) on a weight of 5.0 is a 60% jitter
# so the loss_jitter_multiplier needs to be 0.4
loss_jitter_multiplier = value_map(abs(jitter_list), 0.0, weight_jitter, 1.0, 0.0)
# if items in network_weight list are tensors, convert them to floats
dtype = get_torch_dtype(self.train_config.dtype)
imgs: torch.Tensor = imgs.to(self.device_torch, dtype=dtype)
# split batched images in half so left is negative and right is positive
negative_images, positive_images = torch.chunk(imgs, 2, dim=3)
positive_latents = self.sd.encode_images(positive_images)
negative_latents = self.sd.encode_images(negative_images)
height = positive_images.shape[2]
width = positive_images.shape[3]
batch_size = positive_images.shape[0]
if self.train_config.gradient_checkpointing:
# may get disabled elsewhere
self.sd.unet.enable_gradient_checkpointing()
noise_scheduler = self.sd.noise_scheduler
optimizer = self.optimizer
lr_scheduler = self.lr_scheduler
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
timesteps = torch.randint(0, self.train_config.max_denoising_steps, (1,), device=self.device_torch)
timesteps = timesteps.long()
# get noise
noise_positive = self.sd.get_latent_noise(
pixel_height=height,
pixel_width=width,
batch_size=batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
noise_negative = noise_positive.clone()
# Add noise to the latents according to the noise magnitude at each timestep
# (this is the forward diffusion process)
noisy_positive_latents = noise_scheduler.add_noise(positive_latents, noise_positive, timesteps)
noisy_negative_latents = noise_scheduler.add_noise(negative_latents, noise_negative, timesteps)
noisy_latents = torch.cat([noisy_positive_latents, noisy_negative_latents], dim=0)
noise = torch.cat([noise_positive, noise_negative], dim=0)
timesteps = torch.cat([timesteps, timesteps], dim=0)
network_multiplier = [network_pos_weight * 1.0, network_neg_weight * -1.0]
self.optimizer.zero_grad()
noisy_latents.requires_grad = False
# if training text encoder enable grads, else do context of no grad
with torch.set_grad_enabled(self.train_config.train_text_encoder):
# fix issue with them being tuples sometimes
prompt_list = []
for prompt in prompts:
if isinstance(prompt, tuple):
prompt = prompt[0]
prompt_list.append(prompt)
conditional_embeds = self.sd.encode_prompt(prompt_list).to(self.device_torch, dtype=dtype)
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
# if self.model_config.is_xl:
# # todo also allow for setting this for low ram in general, but sdxl spikes a ton on back prop
# network_multiplier_list = network_multiplier
# noisy_latent_list = torch.chunk(noisy_latents, 2, dim=0)
# noise_list = torch.chunk(noise, 2, dim=0)
# timesteps_list = torch.chunk(timesteps, 2, dim=0)
# conditional_embeds_list = split_prompt_embeds(conditional_embeds)
# else:
network_multiplier_list = [network_multiplier]
noisy_latent_list = [noisy_latents]
noise_list = [noise]
timesteps_list = [timesteps]
conditional_embeds_list = [conditional_embeds]
losses = []
# allow to chunk it out to save vram
for network_multiplier, noisy_latents, noise, timesteps, conditional_embeds in zip(
network_multiplier_list, noisy_latent_list, noise_list, timesteps_list, conditional_embeds_list
):
with self.network:
assert self.network.is_active
self.network.multiplier = network_multiplier
noise_pred = self.sd.predict_noise(
latents=noisy_latents.to(self.device_torch, dtype=dtype),
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
timestep=timesteps,
)
noise = noise.to(self.device_torch, dtype=dtype)
if self.sd.prediction_type == 'v_prediction':
# v-parameterization training
target = noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
else:
target = noise
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
loss = loss.mean([1, 2, 3])
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
# add min_snr_gamma
loss = apply_snr_weight(loss, timesteps, noise_scheduler, self.train_config.min_snr_gamma)
loss = loss.mean() * loss_jitter_multiplier
loss_float = loss.item()
losses.append(loss_float)
# back propagate loss to free ram
loss.backward()
# apply gradients
optimizer.step()
lr_scheduler.step()
# reset network
self.network.multiplier = 1.0
loss_dict = OrderedDict(
{'loss': sum(losses) / len(losses) if len(losses) > 0 else 0.0}
)
return loss_dict
# end hook_train_loop

View File

@@ -0,0 +1,25 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# We make a subclass of Extension
class ImageReferenceSliderTrainer(Extension):
# uid must be unique, it is how the extension is identified
uid = "image_reference_slider_trainer"
# name is the name of the extension for printing
name = "Image Reference Slider 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 .ImageReferenceSliderTrainerProcess import ImageReferenceSliderTrainerProcess
return ImageReferenceSliderTrainerProcess
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
ImageReferenceSliderTrainer
]

View File

@@ -0,0 +1,107 @@
---
job: extension
config:
name: example_name
process:
- type: 'image_reference_slider_trainer'
training_folder: "/mnt/Train/out/LoRA"
device: cuda:0
# for tensorboard logging
log_dir: "/home/jaret/Dev/.tensorboard"
network:
type: "lora"
linear: 8
linear_alpha: 8
train:
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
steps: 5000
lr: 1e-4
train_unet: true
gradient_checkpointing: true
train_text_encoder: true
optimizer: "adamw"
optimizer_params:
weight_decay: 1e-2
lr_scheduler: "constant"
max_denoising_steps: 1000
batch_size: 1
dtype: bf16
xformers: true
skip_first_sample: true
noise_offset: 0.0
model:
name_or_path: "/path/to/model.safetensors"
is_v2: false # for v2 models
is_xl: false # for SDXL models
is_v_pred: false # for v-prediction models (most v2 models)
save:
dtype: float16 # precision to save
save_every: 1000 # save every this many steps
max_step_saves_to_keep: 2 # only affects step counts
sample:
sampler: "ddpm" # must match train.noise_scheduler
sample_every: 100 # sample every this many steps
width: 512
height: 512
prompts:
- "photo of a woman with red hair taking a selfie --m -3"
- "photo of a woman with red hair taking a selfie --m -1"
- "photo of a woman with red hair taking a selfie --m 1"
- "photo of a woman with red hair taking a selfie --m 3"
- "close up photo of a man smiling at the camera, in a tank top --m -3"
- "close up photo of a man smiling at the camera, in a tank top--m -1"
- "close up photo of a man smiling at the camera, in a tank top --m 1"
- "close up photo of a man smiling at the camera, in a tank top --m 3"
- "photo of a blonde woman smiling, barista --m -3"
- "photo of a blonde woman smiling, barista --m -1"
- "photo of a blonde woman smiling, barista --m 1"
- "photo of a blonde woman smiling, barista --m 3"
- "photo of a Christina Hendricks --m -1"
- "photo of a Christina Hendricks --m -1"
- "photo of a Christina Hendricks --m 1"
- "photo of a Christina Hendricks --m 3"
- "photo of a Christina Ricci --m -3"
- "photo of a Christina Ricci --m -1"
- "photo of a Christina Ricci --m 1"
- "photo of a Christina Ricci --m 3"
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
seed: 42
walk_seed: false
guidance_scale: 7
sample_steps: 20
network_multiplier: 1.0
logging:
log_every: 10 # log every this many steps
use_wandb: false # not supported yet
verbose: false
slider:
datasets:
- pair_folder: "/path/to/folder/side/by/side/images"
network_weight: 2.0
target_class: "" # only used as default if caption txt are not present
size: 512
- pair_folder: "/path/to/folder/side/by/side/images"
network_weight: 4.0
target_class: "" # only used as default if caption txt are not present
size: 512
# you can put any information you want here, and it will be saved in the model
# the below is an example. I recommend doing trigger words at a minimum
# in the metadata. The software will include this plus some other information
meta:
name: "[name]" # [name] gets replaced with the name above
description: A short description of your model
trigger_words:
- put
- trigger
- words
- here
version: '0.1'
creator:
name: Your Name
email: your@email.com
website: https://yourwebsite.com
any: All meta data above is arbitrary, it can be whatever you want.

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,249 @@
import os
import random
from collections import OrderedDict
from typing import Union, List
import numpy as np
from diffusers import T2IAdapter, ControlNetModel
import torch.distributed as dist
from torch import nn
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler
from toolkit.clip_vision_adapter import ClipVisionAdapter
from toolkit.data_loader import get_dataloader_datasets
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.stable_diffusion_model import BlankNetwork
from toolkit.train_tools import get_torch_dtype, add_all_snr_to_noise_scheduler
import gc
import torch
from jobs.process import BaseSDTrainProcess
from torchvision import transforms
from diffusers import EMAModel
import math
from toolkit.train_tools import precondition_model_outputs_flow_match
from toolkit.models.unified_training_model import UnifiedTrainingModel
def flush():
torch.cuda.empty_cache()
gc.collect()
adapter_transforms = transforms.Compose([
transforms.ToTensor(),
])
class TrainerV2(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super().__init__(process_id, job, config, **kwargs)
self.assistant_adapter: Union['T2IAdapter', 'ControlNetModel', None]
self.do_prior_prediction = False
self.do_long_prompts = False
self.do_guided_loss = False
self._clip_image_embeds_unconditional: Union[List[str], None] = None
self.negative_prompt_pool: Union[List[str], None] = None
self.batch_negative_prompt: Union[List[str], None] = None
self.scaler = torch.cuda.amp.GradScaler()
self.is_bfloat = self.train_config.dtype == "bfloat16" or self.train_config.dtype == "bf16"
self.do_grad_scale = True
if self.is_fine_tuning:
self.do_grad_scale = False
if self.adapter_config is not None:
if self.adapter_config.train:
self.do_grad_scale = False
if self.train_config.dtype in ["fp16", "float16"]:
# patch the scaler to allow fp16 training
org_unscale_grads = self.scaler._unscale_grads_
def _unscale_grads_replacer(optimizer, inv_scale, found_inf, allow_fp16):
return org_unscale_grads(optimizer, inv_scale, found_inf, True)
self.scaler._unscale_grads_ = _unscale_grads_replacer
self.unified_training_model: UnifiedTrainingModel = None
self.device_ids = list(range(torch.cuda.device_count()))
def before_model_load(self):
pass
def before_dataset_load(self):
self.assistant_adapter = None
# get adapter assistant if one is set
if self.train_config.adapter_assist_name_or_path is not None:
adapter_path = self.train_config.adapter_assist_name_or_path
if self.train_config.adapter_assist_type == "t2i":
# dont name this adapter since we are not training it
self.assistant_adapter = T2IAdapter.from_pretrained(
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype)
).to(self.device_torch)
elif self.train_config.adapter_assist_type == "control_net":
self.assistant_adapter = ControlNetModel.from_pretrained(
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype)
).to(self.device_torch, dtype=get_torch_dtype(self.train_config.dtype))
else:
raise ValueError(f"Unknown adapter assist type {self.train_config.adapter_assist_type}")
self.assistant_adapter.eval()
self.assistant_adapter.requires_grad_(False)
flush()
if self.train_config.train_turbo and self.train_config.show_turbo_outputs:
raise ValueError("Turbo outputs are not supported on MultiGPUSDTrainer")
def hook_before_train_loop(self):
# if self.train_config.do_prior_divergence:
# self.do_prior_prediction = True
# move vae to device if we did not cache latents
if not self.is_latents_cached:
self.sd.vae.eval()
self.sd.vae.to(self.device_torch)
else:
# offload it. Already cached
self.sd.vae.to('cpu')
flush()
add_all_snr_to_noise_scheduler(self.sd.noise_scheduler, self.device_torch)
if self.adapter is not None:
self.adapter.to(self.device_torch)
# check if we have regs and using adapter and caching clip embeddings
has_reg = self.datasets_reg is not None and len(self.datasets_reg) > 0
is_caching_clip_embeddings = self.datasets is not None and any([self.datasets[i].cache_clip_vision_to_disk for i in range(len(self.datasets))])
if has_reg and is_caching_clip_embeddings:
# we need a list of unconditional clip image embeds from other datasets to handle regs
unconditional_clip_image_embeds = []
datasets = get_dataloader_datasets(self.data_loader)
for i in range(len(datasets)):
unconditional_clip_image_embeds += datasets[i].clip_vision_unconditional_cache
if len(unconditional_clip_image_embeds) == 0:
raise ValueError("No unconditional clip image embeds found. This should not happen")
self._clip_image_embeds_unconditional = unconditional_clip_image_embeds
if self.train_config.negative_prompt is not None:
raise ValueError("Negative prompt is not supported on MultiGPUSDTrainer")
# setup the unified training model
self.unified_training_model = UnifiedTrainingModel(
sd=self.sd,
network=self.network,
adapter=self.adapter,
assistant_adapter=self.assistant_adapter,
train_config=self.train_config,
adapter_config=self.adapter_config,
embedding=self.embedding,
timer=self.timer,
trigger_word=self.trigger_word,
gpu_ids=self.device_ids,
)
self.unified_training_model = nn.DataParallel(
self.unified_training_model,
device_ids=self.device_ids
)
self.unified_training_model = self.unified_training_model.to(self.device_torch)
# call parent hook
super().hook_before_train_loop()
# you can expand these in a child class to make customization easier
def preprocess_batch(self, batch: 'DataLoaderBatchDTO'):
return self.unified_training_model.preprocess_batch(batch)
def before_unet_predict(self):
pass
def after_unet_predict(self):
pass
def end_of_training_loop(self):
pass
def hook_train_loop(self, batch: 'DataLoaderBatchDTO'):
self.optimizer.zero_grad(set_to_none=True)
loss = self.unified_training_model(batch)
if torch.isnan(loss):
print("loss is nan")
loss = torch.zeros_like(loss).requires_grad_(True)
if self.network is not None:
network = self.network
else:
network = BlankNetwork()
with (network):
with self.timer('backward'):
# todo we have multiplier seperated. works for now as res are not in same batch, but need to change
# IMPORTANT if gradient checkpointing do not leave with network when doing backward
# it will destroy the gradients. This is because the network is a context manager
# and will change the multipliers back to 0.0 when exiting. They will be
# 0.0 for the backward pass and the gradients will be 0.0
# I spent weeks on fighting this. DON'T DO IT
# with fsdp_overlap_step_with_backward():
# if self.is_bfloat:
# loss.backward()
# else:
if not self.do_grad_scale:
loss.backward()
else:
self.scaler.scale(loss).backward()
if not self.is_grad_accumulation_step:
# fix this for multi params
if self.train_config.optimizer != 'adafactor':
if self.do_grad_scale:
self.scaler.unscale_(self.optimizer)
if isinstance(self.params[0], dict):
for i in range(len(self.params)):
torch.nn.utils.clip_grad_norm_(self.params[i]['params'], self.train_config.max_grad_norm)
else:
torch.nn.utils.clip_grad_norm_(self.params, self.train_config.max_grad_norm)
# only step if we are not accumulating
with self.timer('optimizer_step'):
# self.optimizer.step()
if not self.do_grad_scale:
self.optimizer.step()
else:
self.scaler.step(self.optimizer)
self.scaler.update()
self.optimizer.zero_grad(set_to_none=True)
if self.ema is not None:
with self.timer('ema_update'):
self.ema.update()
else:
# gradient accumulation. Just a place for breakpoint
pass
# TODO Should we only step scheduler on grad step? If so, need to recalculate last step
with self.timer('scheduler_step'):
self.lr_scheduler.step()
if self.embedding is not None:
with self.timer('restore_embeddings'):
# Let's make sure we don't update any embedding weights besides the newly added token
self.embedding.restore_embeddings()
if self.adapter is not None and isinstance(self.adapter, ClipVisionAdapter):
with self.timer('restore_adapter'):
# Let's make sure we don't update any embedding weights besides the newly added token
self.adapter.restore_embeddings()
loss_dict = OrderedDict(
{'loss': loss.item()}
)
self.end_of_training_loop()
return loss_dict

View File

@@ -0,0 +1,47 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# This is for generic training (LoRA, Dreambooth, FineTuning)
class SDTrainerExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "sd_trainer"
# name is the name of the extension for printing
name = "SD 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 .SDTrainer import SDTrainer
return SDTrainer
# This is for generic training (LoRA, Dreambooth, FineTuning)
class MultiGPUSDTrainerExtension(Extension):
# uid must be unique, it is how the extension is identified
uid = "trainer_v2"
# name is the name of the extension for printing
name = "Trainer V2"
# This is where your process class is loaded
# keep your imports in here so they don't slow down the rest of the program
@classmethod
def get_process(cls):
# import your process class here so it is only loaded when needed and return it
from .TrainerV2 import TrainerV2
return TrainerV2
# for backwards compatability
class TextualInversionTrainer(SDTrainerExtension):
uid = "textual_inversion_trainer"
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
SDTrainerExtension, TextualInversionTrainer, MultiGPUSDTrainerExtension
]

View File

@@ -0,0 +1,91 @@
---
job: extension
config:
name: test_v1
process:
- type: 'textual_inversion_trainer'
training_folder: "out/TI"
device: cuda:0
# for tensorboard logging
log_dir: "out/.tensorboard"
embedding:
trigger: "your_trigger_here"
tokens: 12
init_words: "man with short brown hair"
save_format: "safetensors" # 'safetensors' or 'pt'
save:
dtype: float16 # precision to save
save_every: 100 # save every this many steps
max_step_saves_to_keep: 5 # only affects step counts
datasets:
- folder_path: "/path/to/dataset"
caption_ext: "txt"
default_caption: "[trigger]"
buckets: true
resolution: 512
train:
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
steps: 3000
weight_jitter: 0.0
lr: 5e-5
train_unet: false
gradient_checkpointing: true
train_text_encoder: false
optimizer: "adamw"
# optimizer: "prodigy"
optimizer_params:
weight_decay: 1e-2
lr_scheduler: "constant"
max_denoising_steps: 1000
batch_size: 4
dtype: bf16
xformers: true
min_snr_gamma: 5.0
# skip_first_sample: true
noise_offset: 0.0 # not needed for this
model:
# objective reality v2
name_or_path: "https://civitai.com/models/128453?modelVersionId=142465"
is_v2: false # for v2 models
is_xl: false # for SDXL models
is_v_pred: false # for v-prediction models (most v2 models)
sample:
sampler: "ddpm" # must match train.noise_scheduler
sample_every: 100 # sample every this many steps
width: 512
height: 512
prompts:
- "photo of [trigger] laughing"
- "photo of [trigger] smiling"
- "[trigger] close up"
- "dark scene [trigger] frozen"
- "[trigger] nighttime"
- "a painting of [trigger]"
- "a drawing of [trigger]"
- "a cartoon of [trigger]"
- "[trigger] pixar style"
- "[trigger] costume"
neg: ""
seed: 42
walk_seed: false
guidance_scale: 7
sample_steps: 20
network_multiplier: 1.0
logging:
log_every: 10 # log every this many steps
use_wandb: false # not supported yet
verbose: false
# You can put any information you want here, and it will be saved in the model.
# The below is an example, but you can put your grocery list in it if you want.
# It is saved in the model so be aware of that. The software will include this
# plus some other information for you automatically
meta:
# [name] gets replaced with the name above
name: "[name]"
# version: '1.0'
# creator:
# name: Your Name
# email: your@gmail.com
# website: https://your.website

View File

@@ -0,0 +1,533 @@
import copy
import random
from collections import OrderedDict
import os
from contextlib import nullcontext
from typing import Optional, Union, List
from torch.utils.data import ConcatDataset, DataLoader
from toolkit.config_modules import ReferenceDatasetConfig
from toolkit.data_loader import PairedImageDataset
from toolkit.prompt_utils import concat_prompt_embeds, split_prompt_embeds, build_latent_image_batch_for_prompt_pair
from toolkit.stable_diffusion_model import StableDiffusion, PromptEmbeds
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
import gc
from toolkit import train_tools
import torch
from jobs.process import BaseSDTrainProcess
import random
import random
from collections import OrderedDict
from tqdm import tqdm
from toolkit.config_modules import SliderConfig
from toolkit.train_tools import get_torch_dtype, apply_snr_weight
import gc
from toolkit import train_tools
from toolkit.prompt_utils import \
EncodedPromptPair, ACTION_TYPES_SLIDER, \
EncodedAnchor, concat_prompt_pairs, \
concat_anchors, PromptEmbedsCache, encode_prompts_to_cache, build_prompt_pair_batch_from_cache, split_anchors, \
split_prompt_pairs
import torch
def flush():
torch.cuda.empty_cache()
gc.collect()
class UltimateSliderConfig(SliderConfig):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.additional_losses: List[str] = kwargs.get('additional_losses', [])
self.weight_jitter: float = kwargs.get('weight_jitter', 0.0)
self.img_loss_weight: float = kwargs.get('img_loss_weight', 1.0)
self.cfg_loss_weight: float = kwargs.get('cfg_loss_weight', 1.0)
self.datasets: List[ReferenceDatasetConfig] = [ReferenceDatasetConfig(**d) for d in kwargs.get('datasets', [])]
class UltimateSliderTrainerProcess(BaseSDTrainProcess):
sd: StableDiffusion
data_loader: DataLoader = None
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super().__init__(process_id, job, config, **kwargs)
self.prompt_txt_list = None
self.step_num = 0
self.start_step = 0
self.device = self.get_conf('device', self.job.device)
self.device_torch = torch.device(self.device)
self.slider_config = UltimateSliderConfig(**self.get_conf('slider', {}))
self.prompt_cache = PromptEmbedsCache()
self.prompt_pairs: list[EncodedPromptPair] = []
self.anchor_pairs: list[EncodedAnchor] = []
# keep track of prompt chunk size
self.prompt_chunk_size = 1
# store a list of all the prompts from the dataset so we can cache it
self.dataset_prompts = []
self.train_with_dataset = self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0
def load_datasets(self):
if self.data_loader is None and \
self.slider_config.datasets is not None and len(self.slider_config.datasets) > 0:
print(f"Loading datasets")
datasets = []
for dataset in self.slider_config.datasets:
print(f" - Dataset: {dataset.pair_folder}")
config = {
'path': dataset.pair_folder,
'size': dataset.size,
'default_prompt': dataset.target_class,
'network_weight': dataset.network_weight,
'pos_weight': dataset.pos_weight,
'neg_weight': dataset.neg_weight,
'pos_folder': dataset.pos_folder,
'neg_folder': dataset.neg_folder,
}
image_dataset = PairedImageDataset(config)
datasets.append(image_dataset)
# capture all the prompts from it so we can cache the embeds
self.dataset_prompts += image_dataset.get_all_prompts()
concatenated_dataset = ConcatDataset(datasets)
self.data_loader = DataLoader(
concatenated_dataset,
batch_size=self.train_config.batch_size,
shuffle=True,
num_workers=2
)
def before_model_load(self):
pass
def hook_before_train_loop(self):
# load any datasets if they were passed
self.load_datasets()
# read line by line from file
if self.slider_config.prompt_file:
self.print(f"Loading prompt file from {self.slider_config.prompt_file}")
with open(self.slider_config.prompt_file, 'r', encoding='utf-8') as f:
self.prompt_txt_list = f.readlines()
# clean empty lines
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
self.print(f"Found {len(self.prompt_txt_list)} prompts.")
if not self.slider_config.prompt_tensors:
print(f"Prompt tensors not found. Building prompt tensors for {self.train_config.steps} steps.")
# shuffle
random.shuffle(self.prompt_txt_list)
# trim to max steps
self.prompt_txt_list = self.prompt_txt_list[:self.train_config.steps]
# trim list to our max steps
cache = PromptEmbedsCache()
# get encoded latents for our prompts
with torch.no_grad():
# list of neutrals. Can come from file or be empty
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
# build the prompts to cache
prompts_to_cache = []
for neutral in neutral_list:
for target in self.slider_config.targets:
prompt_list = [
f"{target.target_class}", # target_class
f"{target.target_class} {neutral}", # target_class with neutral
f"{target.positive}", # positive_target
f"{target.positive} {neutral}", # positive_target with neutral
f"{target.negative}", # negative_target
f"{target.negative} {neutral}", # negative_target with neutral
f"{neutral}", # neutral
f"{target.positive} {target.negative}", # both targets
f"{target.negative} {target.positive}", # both targets reverse
]
prompts_to_cache += prompt_list
# remove duplicates
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
# trim to max steps if max steps is lower than prompt count
prompts_to_cache = prompts_to_cache[:self.train_config.steps]
if len(self.dataset_prompts) > 0:
# add the prompts from the dataset
prompts_to_cache += self.dataset_prompts
# encode them
cache = encode_prompts_to_cache(
prompt_list=prompts_to_cache,
sd=self.sd,
cache=cache,
prompt_tensor_file=self.slider_config.prompt_tensors
)
prompt_pairs = []
prompt_batches = []
for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
for target in self.slider_config.targets:
prompt_pair_batch = build_prompt_pair_batch_from_cache(
cache=cache,
target=target,
neutral=neutral,
)
if self.slider_config.batch_full_slide:
# concat the prompt pairs
# this allows us to run the entire 4 part process in one shot (for slider)
self.prompt_chunk_size = 4
concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
prompt_pairs += [concat_prompt_pair_batch]
else:
self.prompt_chunk_size = 1
# do them one at a time (probably not necessary after new optimizations)
prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
# move to cpu to save vram
# We don't need text encoder anymore, but keep it on cpu for sampling
# if text encoder is list
if isinstance(self.sd.text_encoder, list):
for encoder in self.sd.text_encoder:
encoder.to("cpu")
else:
self.sd.text_encoder.to("cpu")
self.prompt_cache = cache
self.prompt_pairs = prompt_pairs
# end hook_before_train_loop
# move vae to device so we can encode on the fly
# todo cache latents
self.sd.vae.to(self.device_torch)
self.sd.vae.eval()
self.sd.vae.requires_grad_(False)
if self.train_config.gradient_checkpointing:
# may get disabled elsewhere
self.sd.unet.enable_gradient_checkpointing()
flush()
# end hook_before_train_loop
def hook_train_loop(self, batch):
dtype = get_torch_dtype(self.train_config.dtype)
with torch.no_grad():
### LOOP SETUP ###
noise_scheduler = self.sd.noise_scheduler
optimizer = self.optimizer
lr_scheduler = self.lr_scheduler
### TARGET_PROMPTS ###
# get a random pair
prompt_pair: EncodedPromptPair = self.prompt_pairs[
torch.randint(0, len(self.prompt_pairs), (1,)).item()
]
# move to device and dtype
prompt_pair.to(self.device_torch, dtype=dtype)
### PREP REFERENCE IMAGES ###
imgs, prompts, network_weights = batch
network_pos_weight, network_neg_weight = network_weights
if isinstance(network_pos_weight, torch.Tensor):
network_pos_weight = network_pos_weight.item()
if isinstance(network_neg_weight, torch.Tensor):
network_neg_weight = network_neg_weight.item()
# get an array of random floats between -weight_jitter and weight_jitter
weight_jitter = self.slider_config.weight_jitter
if weight_jitter > 0.0:
jitter_list = random.uniform(-weight_jitter, weight_jitter)
network_pos_weight += jitter_list
network_neg_weight += (jitter_list * -1.0)
# if items in network_weight list are tensors, convert them to floats
imgs: torch.Tensor = imgs.to(self.device_torch, dtype=dtype)
# split batched images in half so left is negative and right is positive
negative_images, positive_images = torch.chunk(imgs, 2, dim=3)
height = positive_images.shape[2]
width = positive_images.shape[3]
batch_size = positive_images.shape[0]
positive_latents = self.sd.encode_images(positive_images)
negative_latents = self.sd.encode_images(negative_images)
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
timesteps = torch.randint(0, self.train_config.max_denoising_steps, (1,), device=self.device_torch)
current_timestep_index = timesteps.item()
current_timestep = noise_scheduler.timesteps[current_timestep_index]
timesteps = timesteps.long()
# get noise
noise_positive = self.sd.get_latent_noise(
pixel_height=height,
pixel_width=width,
batch_size=batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
noise_negative = noise_positive.clone()
# Add noise to the latents according to the noise magnitude at each timestep
# (this is the forward diffusion process)
noisy_positive_latents = noise_scheduler.add_noise(positive_latents, noise_positive, timesteps)
noisy_negative_latents = noise_scheduler.add_noise(negative_latents, noise_negative, timesteps)
### CFG SLIDER TRAINING PREP ###
# get CFG txt latents
noisy_cfg_latents = build_latent_image_batch_for_prompt_pair(
pos_latent=noisy_positive_latents,
neg_latent=noisy_negative_latents,
prompt_pair=prompt_pair,
prompt_chunk_size=self.prompt_chunk_size,
)
noisy_cfg_latents.requires_grad = False
assert not self.network.is_active
# 4.20 GB RAM for 512x512
positive_latents = self.sd.predict_noise(
latents=noisy_cfg_latents,
text_embeddings=train_tools.concat_prompt_embeddings(
prompt_pair.positive_target, # negative prompt
prompt_pair.negative_target, # positive prompt
self.train_config.batch_size,
),
timestep=current_timestep,
guidance_scale=1.0
)
positive_latents.requires_grad = False
neutral_latents = self.sd.predict_noise(
latents=noisy_cfg_latents,
text_embeddings=train_tools.concat_prompt_embeddings(
prompt_pair.positive_target, # negative prompt
prompt_pair.empty_prompt, # positive prompt (normally neutral
self.train_config.batch_size,
),
timestep=current_timestep,
guidance_scale=1.0
)
neutral_latents.requires_grad = False
unconditional_latents = self.sd.predict_noise(
latents=noisy_cfg_latents,
text_embeddings=train_tools.concat_prompt_embeddings(
prompt_pair.positive_target, # negative prompt
prompt_pair.positive_target, # positive prompt
self.train_config.batch_size,
),
timestep=current_timestep,
guidance_scale=1.0
)
unconditional_latents.requires_grad = False
positive_latents_chunks = torch.chunk(positive_latents, self.prompt_chunk_size, dim=0)
neutral_latents_chunks = torch.chunk(neutral_latents, self.prompt_chunk_size, dim=0)
unconditional_latents_chunks = torch.chunk(unconditional_latents, self.prompt_chunk_size, dim=0)
prompt_pair_chunks = split_prompt_pairs(prompt_pair, self.prompt_chunk_size)
noisy_cfg_latents_chunks = torch.chunk(noisy_cfg_latents, self.prompt_chunk_size, dim=0)
assert len(prompt_pair_chunks) == len(noisy_cfg_latents_chunks)
noisy_latents = torch.cat([noisy_positive_latents, noisy_negative_latents], dim=0)
noise = torch.cat([noise_positive, noise_negative], dim=0)
timesteps = torch.cat([timesteps, timesteps], dim=0)
network_multiplier = [network_pos_weight * 1.0, network_neg_weight * -1.0]
flush()
loss_float = None
loss_mirror_float = None
self.optimizer.zero_grad()
noisy_latents.requires_grad = False
# TODO allow both processed to train text encoder, for now, we just to unet and cache all text encodes
# if training text encoder enable grads, else do context of no grad
# with torch.set_grad_enabled(self.train_config.train_text_encoder):
# # text encoding
# embedding_list = []
# # embed the prompts
# for prompt in prompts:
# embedding = self.sd.encode_prompt(prompt).to(self.device_torch, dtype=dtype)
# embedding_list.append(embedding)
# conditional_embeds = concat_prompt_embeds(embedding_list)
# conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
if self.train_with_dataset:
embedding_list = []
with torch.set_grad_enabled(self.train_config.train_text_encoder):
for prompt in prompts:
# get embedding form cache
embedding = self.prompt_cache[prompt]
embedding = embedding.to(self.device_torch, dtype=dtype)
embedding_list.append(embedding)
conditional_embeds = concat_prompt_embeds(embedding_list)
# double up so we can do both sides of the slider
conditional_embeds = concat_prompt_embeds([conditional_embeds, conditional_embeds])
else:
# throw error. Not supported yet
raise Exception("Datasets and targets required for ultimate slider")
if self.model_config.is_xl:
# todo also allow for setting this for low ram in general, but sdxl spikes a ton on back prop
network_multiplier_list = network_multiplier
noisy_latent_list = torch.chunk(noisy_latents, 2, dim=0)
noise_list = torch.chunk(noise, 2, dim=0)
timesteps_list = torch.chunk(timesteps, 2, dim=0)
conditional_embeds_list = split_prompt_embeds(conditional_embeds)
else:
network_multiplier_list = [network_multiplier]
noisy_latent_list = [noisy_latents]
noise_list = [noise]
timesteps_list = [timesteps]
conditional_embeds_list = [conditional_embeds]
## DO REFERENCE IMAGE TRAINING ##
reference_image_losses = []
# allow to chunk it out to save vram
for network_multiplier, noisy_latents, noise, timesteps, conditional_embeds in zip(
network_multiplier_list, noisy_latent_list, noise_list, timesteps_list, conditional_embeds_list
):
with self.network:
assert self.network.is_active
self.network.multiplier = network_multiplier
noise_pred = self.sd.predict_noise(
latents=noisy_latents.to(self.device_torch, dtype=dtype),
conditional_embeddings=conditional_embeds.to(self.device_torch, dtype=dtype),
timestep=timesteps,
)
noise = noise.to(self.device_torch, dtype=dtype)
if self.sd.prediction_type == 'v_prediction':
# v-parameterization training
target = noise_scheduler.get_velocity(noisy_latents, noise, timesteps)
else:
target = noise
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
loss = loss.mean([1, 2, 3])
# todo add snr gamma here
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
# add min_snr_gamma
loss = apply_snr_weight(loss, timesteps, noise_scheduler, self.train_config.min_snr_gamma)
loss = loss.mean()
loss = loss * self.slider_config.img_loss_weight
loss_slide_float = loss.item()
loss_float = loss.item()
reference_image_losses.append(loss_float)
# back propagate loss to free ram
loss.backward()
flush()
## DO CFG SLIDER TRAINING ##
cfg_loss_list = []
with self.network:
assert self.network.is_active
for prompt_pair_chunk, \
noisy_cfg_latent_chunk, \
positive_latents_chunk, \
neutral_latents_chunk, \
unconditional_latents_chunk \
in zip(
prompt_pair_chunks,
noisy_cfg_latents_chunks,
positive_latents_chunks,
neutral_latents_chunks,
unconditional_latents_chunks,
):
self.network.multiplier = prompt_pair_chunk.multiplier_list
target_latents = self.sd.predict_noise(
latents=noisy_cfg_latent_chunk,
text_embeddings=train_tools.concat_prompt_embeddings(
prompt_pair_chunk.positive_target, # negative prompt
prompt_pair_chunk.target_class, # positive prompt
self.train_config.batch_size,
),
timestep=current_timestep,
guidance_scale=1.0
)
guidance_scale = 1.0
offset = guidance_scale * (positive_latents_chunk - unconditional_latents_chunk)
# make offset multiplier based on actions
offset_multiplier_list = []
for action in prompt_pair_chunk.action_list:
if action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE:
offset_multiplier_list += [-1.0]
elif action == ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE:
offset_multiplier_list += [1.0]
offset_multiplier = torch.tensor(offset_multiplier_list).to(offset.device, dtype=offset.dtype)
# make offset multiplier match rank of offset
offset_multiplier = offset_multiplier.view(offset.shape[0], 1, 1, 1)
offset *= offset_multiplier
offset_neutral = neutral_latents_chunk
# offsets are already adjusted on a per-batch basis
offset_neutral += offset
# 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")
loss = loss.mean([1, 2, 3])
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
# match batch size
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
# add min_snr_gamma
loss = apply_snr_weight(loss, timesteps_index_list, noise_scheduler,
self.train_config.min_snr_gamma)
loss = loss.mean() * prompt_pair_chunk.weight * self.slider_config.cfg_loss_weight
loss.backward()
cfg_loss_list.append(loss.item())
del target_latents
del offset_neutral
del loss
flush()
# apply gradients
optimizer.step()
lr_scheduler.step()
# reset network
self.network.multiplier = 1.0
reference_image_loss = sum(reference_image_losses) / len(reference_image_losses) if len(
reference_image_losses) > 0 else 0.0
cfg_loss = sum(cfg_loss_list) / len(cfg_loss_list) if len(cfg_loss_list) > 0 else 0.0
loss_dict = OrderedDict({
'loss/img': reference_image_loss,
'loss/cfg': cfg_loss,
})
return loss_dict
# end hook_train_loop

View File

@@ -0,0 +1,25 @@
# This is an example extension for custom training. It is great for experimenting with new ideas.
from toolkit.extension import Extension
# We make a subclass of Extension
class UltimateSliderTrainer(Extension):
# uid must be unique, it is how the extension is identified
uid = "ultimate_slider_trainer"
# name is the name of the extension for printing
name = "Ultimate Slider 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 .UltimateSliderTrainerProcess import UltimateSliderTrainerProcess
return UltimateSliderTrainerProcess
AI_TOOLKIT_EXTENSIONS = [
# you can put a list of extensions here
UltimateSliderTrainer
]

View File

@@ -0,0 +1,107 @@
---
job: extension
config:
name: example_name
process:
- type: 'image_reference_slider_trainer'
training_folder: "/mnt/Train/out/LoRA"
device: cuda:0
# for tensorboard logging
log_dir: "/home/jaret/Dev/.tensorboard"
network:
type: "lora"
linear: 8
linear_alpha: 8
train:
noise_scheduler: "ddpm" # or "ddpm", "lms", "euler_a"
steps: 5000
lr: 1e-4
train_unet: true
gradient_checkpointing: true
train_text_encoder: true
optimizer: "adamw"
optimizer_params:
weight_decay: 1e-2
lr_scheduler: "constant"
max_denoising_steps: 1000
batch_size: 1
dtype: bf16
xformers: true
skip_first_sample: true
noise_offset: 0.0
model:
name_or_path: "/path/to/model.safetensors"
is_v2: false # for v2 models
is_xl: false # for SDXL models
is_v_pred: false # for v-prediction models (most v2 models)
save:
dtype: float16 # precision to save
save_every: 1000 # save every this many steps
max_step_saves_to_keep: 2 # only affects step counts
sample:
sampler: "ddpm" # must match train.noise_scheduler
sample_every: 100 # sample every this many steps
width: 512
height: 512
prompts:
- "photo of a woman with red hair taking a selfie --m -3"
- "photo of a woman with red hair taking a selfie --m -1"
- "photo of a woman with red hair taking a selfie --m 1"
- "photo of a woman with red hair taking a selfie --m 3"
- "close up photo of a man smiling at the camera, in a tank top --m -3"
- "close up photo of a man smiling at the camera, in a tank top--m -1"
- "close up photo of a man smiling at the camera, in a tank top --m 1"
- "close up photo of a man smiling at the camera, in a tank top --m 3"
- "photo of a blonde woman smiling, barista --m -3"
- "photo of a blonde woman smiling, barista --m -1"
- "photo of a blonde woman smiling, barista --m 1"
- "photo of a blonde woman smiling, barista --m 3"
- "photo of a Christina Hendricks --m -1"
- "photo of a Christina Hendricks --m -1"
- "photo of a Christina Hendricks --m 1"
- "photo of a Christina Hendricks --m 3"
- "photo of a Christina Ricci --m -3"
- "photo of a Christina Ricci --m -1"
- "photo of a Christina Ricci --m 1"
- "photo of a Christina Ricci --m 3"
neg: "cartoon, fake, drawing, illustration, cgi, animated, anime"
seed: 42
walk_seed: false
guidance_scale: 7
sample_steps: 20
network_multiplier: 1.0
logging:
log_every: 10 # log every this many steps
use_wandb: false # not supported yet
verbose: false
slider:
datasets:
- pair_folder: "/path/to/folder/side/by/side/images"
network_weight: 2.0
target_class: "" # only used as default if caption txt are not present
size: 512
- pair_folder: "/path/to/folder/side/by/side/images"
network_weight: 4.0
target_class: "" # only used as default if caption txt are not present
size: 512
# you can put any information you want here, and it will be saved in the model
# the below is an example. I recommend doing trigger words at a minimum
# in the metadata. The software will include this plus some other information
meta:
name: "[name]" # [name] gets replaced with the name above
description: A short description of your model
trigger_words:
- put
- trigger
- words
- here
version: '0.1'
creator:
name: Your Name
email: your@email.com
website: https://yourwebsite.com
any: All meta data above is arbitrary, it can be whatever you want.

View File

@@ -3,6 +3,6 @@ from collections import OrderedDict
v = OrderedDict()
v["name"] = "ai-toolkit"
v["repo"] = "https://github.com/ostris/ai-toolkit"
v["version"] = "0.0.1"
v["version"] = "0.1.0"
software_meta = v

View File

@@ -6,19 +6,16 @@ from jobs.process import BaseProcess
class BaseJob:
config: OrderedDict
job: str
name: str
meta: OrderedDict
process: List[BaseProcess]
def __init__(self, config: OrderedDict):
if not config:
raise ValueError('config is required')
self.process: List[BaseProcess]
self.config = config['config']
self.raw_config = config
self.job = config['job']
self.torch_profiler = self.get_conf('torch_profiler', False)
self.name = self.get_conf('name', required=True)
if 'meta' in config:
self.meta = config['meta']
@@ -60,7 +57,11 @@ class BaseJob:
# check if dict key is process type
if process['type'] in process_dict:
ProcessClass = getattr(module, process_dict[process['type']])
if isinstance(process_dict[process['type']], str):
ProcessClass = getattr(module, process_dict[process['type']])
else:
# it is the class
ProcessClass = process_dict[process['type']]
self.process.append(ProcessClass(i, self, process))
else:
raise ValueError(f'config file is invalid. Unknown process type: {process["type"]}')

22
jobs/ExtensionJob.py Normal file
View File

@@ -0,0 +1,22 @@
import os
from collections import OrderedDict
from jobs import BaseJob
from toolkit.extension import get_all_extensions_process_dict
from toolkit.paths import CONFIG_ROOT
class ExtensionJob(BaseJob):
def __init__(self, config: OrderedDict):
super().__init__(config)
self.device = self.get_conf('device', 'cpu')
self.process_dict = get_all_extensions_process_dict()
self.load_processes(self.process_dict)
def run(self):
super().run()
print("")
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
for process in self.process:
process.run()

31
jobs/GenerateJob.py Normal file
View File

@@ -0,0 +1,31 @@
from jobs import BaseJob
from collections import OrderedDict
from typing import List
from jobs.process import GenerateProcess
from toolkit.paths import REPOS_ROOT
import sys
sys.path.append(REPOS_ROOT)
process_dict = {
'to_folder': 'GenerateProcess',
}
class GenerateJob(BaseJob):
def __init__(self, config: OrderedDict):
super().__init__(config)
self.device = self.get_conf('device', 'cpu')
# loads the processes from the config
self.load_processes(process_dict)
def run(self):
super().run()
print("")
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
for process in self.process:
process.run()

28
jobs/ModJob.py Normal file
View File

@@ -0,0 +1,28 @@
import os
from collections import OrderedDict
from jobs import BaseJob
from toolkit.metadata import get_meta_for_safetensors
from toolkit.train_tools import get_torch_dtype
process_dict = {
'rescale_lora': 'ModRescaleLoraProcess',
}
class ModJob(BaseJob):
def __init__(self, config: OrderedDict):
super().__init__(config)
self.device = self.get_conf('device', 'cpu')
# loads the processes from the config
self.load_processes(process_dict)
def run(self):
super().run()
print("")
print(f"Running {len(self.process)} process{'' if len(self.process) == 1 else 'es'}")
for process in self.process:
process.run()

View File

@@ -17,13 +17,15 @@ sys.path.append(REPOS_ROOT)
process_dict = {
'vae': 'TrainVAEProcess',
'slider': 'TrainSliderProcess',
'slider_old': 'TrainSliderProcessOld',
'lora_hack': 'TrainLoRAHack',
'rescale_sd': 'TrainSDRescaleProcess',
'esrgan': 'TrainESRGANProcess',
'reference': 'TrainReferenceProcess',
}
class TrainJob(BaseJob):
process: List[BaseExtractProcess]
def __init__(self, config: OrderedDict):
super().__init__(config)
@@ -34,18 +36,9 @@ class TrainJob(BaseJob):
# self.mixed_precision = self.get_conf('mixed_precision', False) # fp16
self.log_dir = self.get_conf('log_dir', None)
self.writer = None
self.setup_tensorboard()
# loads the processes from the config
self.load_processes(process_dict)
def save_training_config(self):
timestamp = datetime.now().strftime('%Y%m%d-%H%M%S')
os.makedirs(self.training_folder, exist_ok=True)
save_dif = os.path.join(self.training_folder, f'run_config_{timestamp}.yaml')
with open(save_dif, 'w') as f:
yaml.dump(self.raw_config, f)
def run(self):
super().run()
@@ -54,12 +47,3 @@ class TrainJob(BaseJob):
for process in self.process:
process.run()
def setup_tensorboard(self):
if self.log_dir:
from torch.utils.tensorboard import SummaryWriter
now = datetime.now()
time_str = now.strftime('%Y%m%d-%H%M%S')
summary_name = f"{self.name}_{time_str}"
summary_dir = os.path.join(self.log_dir, summary_name)
self.writer = SummaryWriter(summary_dir)

View File

@@ -2,3 +2,6 @@ from .BaseJob import BaseJob
from .ExtractJob import ExtractJob
from .TrainJob import TrainJob
from .MergeJob import MergeJob
from .ModJob import ModJob
from .GenerateJob import GenerateJob
from .ExtensionJob import ExtensionJob

View File

@@ -0,0 +1,19 @@
from collections import OrderedDict
from typing import ForwardRef
from jobs.process.BaseProcess import BaseProcess
class BaseExtensionProcess(BaseProcess):
def __init__(
self,
process_id: int,
job,
config: OrderedDict
):
super().__init__(process_id, job, config)
self.process_id: int
self.config: OrderedDict
self.progress_bar: ForwardRef('tqdm') = None
def run(self):
super().run()

View File

@@ -12,11 +12,6 @@ from toolkit.train_tools import get_torch_dtype
class BaseExtractProcess(BaseProcess):
process_id: int
config: OrderedDict
output_folder: str
output_filename: str
output_path: str
def __init__(
self,
@@ -25,6 +20,10 @@ class BaseExtractProcess(BaseProcess):
config: OrderedDict
):
super().__init__(process_id, job, config)
self.config: OrderedDict
self.output_folder: str
self.output_filename: str
self.output_path: str
self.process_id = process_id
self.job = job
self.config = config

View File

@@ -9,8 +9,6 @@ from toolkit.train_tools import get_torch_dtype
class BaseMergeProcess(BaseProcess):
process_id: int
config: OrderedDict
def __init__(
self,
@@ -19,6 +17,8 @@ class BaseMergeProcess(BaseProcess):
config: OrderedDict
):
super().__init__(process_id, job, config)
self.process_id: int
self.config: OrderedDict
self.output_path = self.get_conf('output_path', required=True)
self.dtype = self.get_conf('dtype', self.job.dtype)
self.torch_dtype = get_torch_dtype(self.dtype)

View File

@@ -1,11 +1,11 @@
import copy
import json
from collections import OrderedDict
from typing import ForwardRef
from toolkit.timer import Timer
class BaseProcess:
meta: OrderedDict
class BaseProcess(object):
def __init__(
self,
@@ -14,9 +14,15 @@ class BaseProcess:
config: OrderedDict
):
self.process_id = process_id
self.meta: OrderedDict
self.job = job
self.config = config
self.raw_process_config = config
self.name = self.get_conf('name', self.job.name)
self.meta = copy.deepcopy(self.job.meta)
self.timer: Timer = Timer(f'{self.name} Timer')
self.performance_log_every = self.get_conf('performance_log_every', 0)
print(json.dumps(self.config, indent=4))
def get_conf(self, key, default=None, required=False, as_type=None):

File diff suppressed because it is too large Load Diff

View File

@@ -1,14 +1,21 @@
import random
from datetime import datetime
import os
from collections import OrderedDict
from typing import ForwardRef
from typing import TYPE_CHECKING, Union
import torch
import yaml
from jobs.process.BaseProcess import BaseProcess
if TYPE_CHECKING:
from jobs import TrainJob, BaseJob, ExtensionJob
from torch.utils.tensorboard import SummaryWriter
from tqdm import tqdm
class BaseTrainProcess(BaseProcess):
process_id: int
config: OrderedDict
progress_bar: ForwardRef('tqdm') = None
def __init__(
self,
@@ -17,12 +24,30 @@ class BaseTrainProcess(BaseProcess):
config: OrderedDict
):
super().__init__(process_id, job, config)
self.process_id: int
self.config: OrderedDict
self.writer: 'SummaryWriter'
self.job: Union['TrainJob', 'BaseJob', 'ExtensionJob']
self.progress_bar: 'tqdm' = None
self.training_seed = self.get_conf('training_seed', self.job.training_seed if hasattr(self.job, 'training_seed') else None)
# if training seed is set, use it
if self.training_seed is not None:
torch.manual_seed(self.training_seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(self.training_seed)
random.seed(self.training_seed)
self.progress_bar = None
self.writer = self.job.writer
self.training_folder = self.get_conf('training_folder', self.job.training_folder)
self.save_root = os.path.join(self.training_folder, self.job.name)
self.writer = None
self.training_folder = self.get_conf('training_folder',
self.job.training_folder if hasattr(self.job, 'training_folder') else None)
self.save_root = os.path.join(self.training_folder, self.name)
self.step = 0
self.first_step = 0
self.log_dir = self.get_conf('log_dir', self.job.log_dir if hasattr(self.job, 'log_dir') else None)
self.setup_tensorboard()
self.save_training_config()
def run(self):
super().run()
@@ -37,3 +62,18 @@ class BaseTrainProcess(BaseProcess):
self.progress_bar.update()
else:
print(*args)
def setup_tensorboard(self):
if self.log_dir:
from torch.utils.tensorboard import SummaryWriter
now = datetime.now()
time_str = now.strftime('%Y%m%d-%H%M%S')
summary_name = f"{self.name}_{time_str}"
summary_dir = os.path.join(self.log_dir, summary_name)
self.writer = SummaryWriter(summary_dir)
def save_training_config(self):
os.makedirs(self.save_root, exist_ok=True)
save_dif = os.path.join(self.save_root, f'config.yaml')
with open(save_dif, 'w') as f:
yaml.dump(self.job.raw_config, f)

View File

@@ -0,0 +1,144 @@
import gc
import os
from collections import OrderedDict
from typing import ForwardRef, List, Optional, Union
import torch
from safetensors.torch import save_file, load_file
from jobs.process.BaseProcess import BaseProcess
from toolkit.config_modules import ModelConfig, GenerateImageConfig
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_model_hash_to_meta, \
add_base_model_info_to_meta
from toolkit.stable_diffusion_model import StableDiffusion
from toolkit.train_tools import get_torch_dtype
import random
class GenerateConfig:
def __init__(self, **kwargs):
self.prompts: List[str]
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)
self.sample_steps = kwargs.get('sample_steps', 20)
self.prompt_2 = kwargs.get('prompt_2', None)
self.neg_2 = kwargs.get('neg_2', None)
self.prompts = kwargs.get('prompts', None)
self.guidance_rescale = kwargs.get('guidance_rescale', 0.0)
self.compile = kwargs.get('compile', False)
self.ext = kwargs.get('ext', 'png')
self.prompt_file = kwargs.get('prompt_file', False)
self.prompts_in_file = self.prompts
if self.prompts is None:
raise ValueError("Prompts must be set")
if isinstance(self.prompts, str):
if os.path.exists(self.prompts):
with open(self.prompts, 'r', encoding='utf-8') as f:
self.prompts_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)
class GenerateProcess(BaseProcess):
process_id: int
config: OrderedDict
progress_bar: ForwardRef('tqdm') = None
sd: StableDiffusion
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.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(
device=self.device,
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):
with torch.no_grad():
super().run()
print("Loading model...")
self.sd.load_model()
self.sd.pipeline.to(self.device, self.torch_dtype)
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(f"Generating {len(self.generate_config.prompts)} images")
# build prompt image configs
prompt_image_configs = []
for prompt in self.generate_config.prompts:
width = self.generate_config.width
height = self.generate_config.height
prompt = self.clean_prompt(prompt)
if self.generate_config.size_list is not None:
# randomly select a size
width, height = random.choice(self.generate_config.size_list)
prompt_image_configs.append(GenerateImageConfig(
prompt=prompt,
prompt_2=self.generate_config.prompt_2,
width=width,
height=height,
num_inference_steps=self.generate_config.sample_steps,
guidance_scale=self.generate_config.guidance_scale,
negative_prompt=self.generate_config.neg,
negative_prompt_2=self.generate_config.neg_2,
seed=self.generate_config.seed,
guidance_rescale=self.generate_config.guidance_rescale,
output_ext=self.generate_config.ext,
output_folder=self.output_folder,
add_prompt_file=self.generate_config.prompt_file
))
# generate images
self.sd.generate_images(prompt_image_configs, sampler=self.generate_config.sampler)
print("Done generating images")
# cleanup
del self.sd
gc.collect()
torch.cuda.empty_cache()

View File

@@ -0,0 +1,104 @@
import gc
import os
from collections import OrderedDict
from typing import ForwardRef
import torch
from safetensors.torch import save_file, load_file
from jobs.process.BaseProcess import BaseProcess
from toolkit.metadata import get_meta_for_safetensors, load_metadata_from_safetensors, add_model_hash_to_meta, \
add_base_model_info_to_meta
from toolkit.train_tools import get_torch_dtype
class ModRescaleLoraProcess(BaseProcess):
process_id: int
config: OrderedDict
progress_bar: ForwardRef('tqdm') = None
def __init__(
self,
process_id: int,
job,
config: OrderedDict
):
super().__init__(process_id, job, config)
self.process_id: int
self.config: OrderedDict
self.progress_bar: ForwardRef('tqdm') = None
self.input_path = self.get_conf('input_path', required=True)
self.output_path = self.get_conf('output_path', required=True)
self.replace_meta = self.get_conf('replace_meta', default=False)
self.save_dtype = self.get_conf('save_dtype', default='fp16', as_type=get_torch_dtype)
self.current_weight = self.get_conf('current_weight', required=True, as_type=float)
self.target_weight = self.get_conf('target_weight', required=True, as_type=float)
self.scale_target = self.get_conf('scale_target', default='up_down') # alpha or up_down
self.is_xl = self.get_conf('is_xl', default=False, as_type=bool)
self.is_v2 = self.get_conf('is_v2', default=False, as_type=bool)
self.progress_bar = None
def run(self):
super().run()
source_state_dict = load_file(self.input_path)
source_meta = load_metadata_from_safetensors(self.input_path)
if self.replace_meta:
self.meta.update(
add_base_model_info_to_meta(
self.meta,
is_xl=self.is_xl,
is_v2=self.is_v2,
)
)
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
else:
save_meta = get_meta_for_safetensors(source_meta, self.job.name, add_software_info=False)
# save
os.makedirs(os.path.dirname(self.output_path), exist_ok=True)
new_state_dict = OrderedDict()
for key in list(source_state_dict.keys()):
v = source_state_dict[key]
v = v.detach().clone().to("cpu").to(get_torch_dtype('fp32'))
# all loras have an alpha, up weight and down weight
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.alpha",
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.lora_down.weight",
# - "lora_te_text_model_encoder_layers_0_mlp_fc1.lora_up.weight",
# we can rescale by adjusting the alpha or the up weights, or the up and down weights
# I assume doing both up and down would be best all around, but I'm not sure
# some locons also have mid weights, we will leave those alone for now, will work without them
# when adjusting alpha, it is used to calculate the multiplier in a lora module
# - scale = alpha / lora_dim
# - output = layer_out + lora_up_out * multiplier * scale
total_module_scale = torch.tensor(self.current_weight / self.target_weight) \
.to("cpu", dtype=get_torch_dtype('fp32'))
num_modules_layers = 2 # up and down
up_down_scale = torch.pow(total_module_scale, 1.0 / num_modules_layers) \
.to("cpu", dtype=get_torch_dtype('fp32'))
# only update alpha
if self.scale_target == 'alpha' and key.endswith('.alpha'):
v = v * total_module_scale
if self.scale_target == 'up_down' and key.endswith('.lora_up.weight') or key.endswith('.lora_down.weight'):
# would it be better to adjust the up weights for fp16 precision? Doing both should reduce chance of NaN
v = v * up_down_scale
v = v.detach().clone().to("cpu").to(self.save_dtype)
new_state_dict[key] = v
save_meta = add_model_hash_to_meta(new_state_dict, save_meta)
save_file(new_state_dict, self.output_path, save_meta)
# cleanup incase there are other jobs
del new_state_dict
del source_state_dict
del source_meta
torch.cuda.empty_cache()
gc.collect()
print(f"Saved to {self.output_path}")

View File

@@ -0,0 +1,657 @@
import copy
import glob
import os
import time
from collections import OrderedDict
from typing import List, Optional
from PIL import Image
from PIL.ImageOps import exif_transpose
from toolkit.basic import flush
from toolkit.models.RRDB import RRDBNet as ESRGAN, esrgan_safetensors_keys
from safetensors.torch import save_file, load_file
from torch.utils.data import DataLoader, ConcatDataset
import torch
from torch import nn
from torchvision.transforms import transforms
from jobs.process import BaseTrainProcess
from toolkit.data_loader import AugmentedImageDataset
from toolkit.esrgan_utils import convert_state_dict_to_basicsr, convert_basicsr_state_dict_to_save_format
from toolkit.losses import ComparativeTotalVariation, get_gradient_penalty, PatternLoss
from toolkit.metadata import get_meta_for_safetensors
from toolkit.optimizer import get_optimizer
from toolkit.style import get_style_model_and_losses
from toolkit.train_tools import get_torch_dtype
from diffusers import AutoencoderKL
from tqdm import tqdm
import time
import numpy as np
from .models.vgg19_critic import Critic
IMAGE_TRANSFORMS = transforms.Compose(
[
transforms.ToTensor(),
# transforms.Normalize([0.5], [0.5]),
]
)
class TrainESRGANProcess(BaseTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.data_loader = None
self.model: ESRGAN = None
self.device = self.get_conf('device', self.job.device)
self.pretrained_path = self.get_conf('pretrained_path', 'None')
self.datasets_objects = self.get_conf('datasets', required=True)
self.batch_size = self.get_conf('batch_size', 1, as_type=int)
self.resolution = self.get_conf('resolution', 256, as_type=int)
self.learning_rate = self.get_conf('learning_rate', 1e-6, as_type=float)
self.sample_every = self.get_conf('sample_every', None)
self.optimizer_type = self.get_conf('optimizer', 'adam')
self.epochs = self.get_conf('epochs', None, as_type=int)
self.max_steps = self.get_conf('max_steps', None, as_type=int)
self.save_every = self.get_conf('save_every', None)
self.upscale_sample = self.get_conf('upscale_sample', 4)
self.dtype = self.get_conf('dtype', 'float32')
self.sample_sources = self.get_conf('sample_sources', None)
self.log_every = self.get_conf('log_every', 100, as_type=int)
self.style_weight = self.get_conf('style_weight', 0, as_type=float)
self.content_weight = self.get_conf('content_weight', 0, as_type=float)
self.mse_weight = self.get_conf('mse_weight', 1e0, as_type=float)
self.zoom = self.get_conf('zoom', 4, as_type=int)
self.tv_weight = self.get_conf('tv_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', {})
self.augmentations = self.get_conf('augmentations', {})
self.torch_dtype = get_torch_dtype(self.dtype)
if self.torch_dtype == torch.bfloat16:
self.esrgan_dtype = torch.float32
else:
self.esrgan_dtype = torch.float32
self.vgg_19 = None
self.style_weight_scalers = []
self.content_weight_scalers = []
# throw error if zoom if not divisible by 2
if self.zoom % 2 != 0:
raise ValueError('zoom must be divisible by 2')
self.step_num = 0
self.epoch_num = 0
self.use_critic = self.get_conf('use_critic', False, as_type=bool)
self.critic = None
if self.use_critic:
self.critic = Critic(
device=self.device,
dtype=self.dtype,
process=self,
**self.get_conf('critic', {}) # pass any other params
)
if self.sample_every is not None and self.sample_sources is None:
raise ValueError('sample_every is specified but sample_sources is not')
if self.epochs is None and self.max_steps is None:
raise ValueError('epochs or max_steps must be specified')
self.data_loaders = []
# check datasets
assert isinstance(self.datasets_objects, list)
for dataset in self.datasets_objects:
if 'path' not in dataset:
raise ValueError('dataset must have a path')
# check if is dir
if not os.path.isdir(dataset['path']):
raise ValueError(f"dataset path does is not a directory: {dataset['path']}")
# make training folder
if not os.path.exists(self.save_root):
os.makedirs(self.save_root, exist_ok=True)
self._pattern_loss = None
# build augmentation transforms
aug_transforms = []
def update_training_metadata(self):
self.add_meta(OrderedDict({"training_info": self.get_training_info()}))
def get_training_info(self):
info = OrderedDict({
'step': self.step_num,
'epoch': self.epoch_num,
})
return info
def load_datasets(self):
if self.data_loader is None:
print(f"Loading datasets")
datasets = []
for dataset in self.datasets_objects:
print(f" - Dataset: {dataset['path']}")
ds = copy.copy(dataset)
ds['resolution'] = self.resolution
if 'augmentations' not in ds:
ds['augmentations'] = self.augmentations
# add the resize down augmentation
ds['augmentations'] = [{
'method': 'Resize',
'params': {
'width': int(self.resolution // self.zoom),
'height': int(self.resolution // self.zoom),
# downscale interpolation, string will be evaluated
'interpolation': 'cv2.INTER_AREA'
}
}] + ds['augmentations']
image_dataset = AugmentedImageDataset(ds)
datasets.append(image_dataset)
concatenated_dataset = ConcatDataset(datasets)
self.data_loader = DataLoader(
concatenated_dataset,
batch_size=self.batch_size,
shuffle=True,
num_workers=6
)
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(
single_target=True,
device=self.device,
output_layer_name='pool_4',
dtype=self.torch_dtype
)
self.vgg_19.to(self.device, dtype=self.torch_dtype)
self.vgg_19.requires_grad_(False)
# we run random noise through first to get layer scalers to normalize the loss per layer
# bs of 2 because we run pred and target through stacked
noise = torch.randn((2, 3, self.resolution, self.resolution), device=self.device, dtype=self.torch_dtype)
self.vgg_19(noise)
for style_loss in self.style_losses:
# get a scaler to normalize to 1
scaler = 1 / torch.mean(style_loss.loss).item()
self.style_weight_scalers.append(scaler)
for content_loss in self.content_losses:
# get a scaler to normalize to 1
scaler = 1 / torch.mean(content_loss.loss).item()
# if is nan, set to 1
if scaler != scaler:
scaler = 1
print(f"Warning: content loss scaler is nan, setting to 1")
self.content_weight_scalers.append(scaler)
self.print(f"Style weight scalers: {self.style_weight_scalers}")
self.print(f"Content weight scalers: {self.content_weight_scalers}")
def get_style_loss(self):
if self.style_weight > 0:
# scale all losses with loss scalers
loss = torch.sum(
torch.stack([loss.loss * scaler for loss, scaler in zip(self.style_losses, self.style_weight_scalers)]))
return loss
else:
return torch.tensor(0.0, device=self.device)
def get_content_loss(self):
if self.content_weight > 0:
# scale all losses with loss scalers
loss = torch.sum(torch.stack(
[loss.loss * scaler for loss, scaler in zip(self.content_losses, self.content_weight_scalers)]))
return loss
else:
return torch.tensor(0.0, device=self.device)
def get_mse_loss(self, pred, target):
if self.mse_weight > 0:
loss_fn = nn.MSELoss()
loss = loss_fn(pred, target)
return loss
else:
return torch.tensor(0.0, device=self.device)
def get_tv_loss(self, pred, target):
if self.tv_weight > 0:
get_tv_loss = ComparativeTotalVariation()
loss = get_tv_loss(pred, target)
return loss
else:
return torch.tensor(0.0, device=self.device)
def get_pattern_loss(self, pred, target):
if self._pattern_loss is None:
self._pattern_loss = PatternLoss(
pattern_size=self.zoom,
dtype=self.torch_dtype
).to(self.device, dtype=self.torch_dtype)
self._pattern_loss = self._pattern_loss.to(self.device, dtype=self.torch_dtype)
loss = torch.mean(self._pattern_loss(pred, target))
return loss
def save(self, step=None):
if not os.path.exists(self.save_root):
os.makedirs(self.save_root, exist_ok=True)
step_num = ''
if step is not None:
# zeropad 9 digits
step_num = f"_{str(step).zfill(9)}"
self.update_training_metadata()
# filename = f'{self.job.name}{step_num}.safetensors'
filename = f'{self.job.name}{step_num}.pth'
# prepare meta
save_meta = get_meta_for_safetensors(self.meta, self.job.name)
# state_dict = self.model.state_dict()
# state has the original state dict keys so we can save what we started from
save_state_dict = self.model.state_dict()
for key in list(save_state_dict.keys()):
v = save_state_dict[key]
v = v.detach().clone().to("cpu").to(torch.float32)
save_state_dict[key] = v
# most things wont use safetensors, save as torch
# save_file(save_state_dict, os.path.join(self.save_root, filename), save_meta)
torch.save(save_state_dict, os.path.join(self.save_root, filename))
self.print(f"Saved to {os.path.join(self.save_root, filename)}")
if self.use_critic:
self.critic.save(step)
def sample(self, step=None, batch: Optional[List[torch.Tensor]] = None):
sample_folder = os.path.join(self.save_root, 'samples')
if not os.path.exists(sample_folder):
os.makedirs(sample_folder, exist_ok=True)
batch_sample_folder = os.path.join(self.save_root, 'samples_batch')
batch_targets = None
batch_inputs = None
if batch is not None and not os.path.exists(batch_sample_folder):
os.makedirs(batch_sample_folder, exist_ok=True)
self.model.eval()
def process_and_save(img, target_img, save_path):
img = img.to(self.device, dtype=self.esrgan_dtype)
output = self.model(img)
# output = (output / 2 + 0.5).clamp(0, 1)
output = output.clamp(0, 1)
img = img.clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
output = output.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
img = img.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
# convert to pillow image
output = Image.fromarray((output * 255).astype(np.uint8))
img = Image.fromarray((img * 255).astype(np.uint8))
if isinstance(target_img, torch.Tensor):
# convert to pil
target_img = target_img.cpu().permute(0, 2, 3, 1).squeeze(0).float().numpy()
target_img = Image.fromarray((target_img * 255).astype(np.uint8))
# upscale to size * self.upscale_sample while maintaining pixels
output = output.resize(
(self.resolution * self.upscale_sample, self.resolution * self.upscale_sample),
resample=Image.NEAREST
)
img = img.resize(
(self.resolution * self.upscale_sample, self.resolution * self.upscale_sample),
resample=Image.NEAREST
)
width, height = output.size
# stack input image and decoded image
target_image = target_img.resize((width, height))
output = output.resize((width, height))
img = img.resize((width, height))
output_img = Image.new('RGB', (width * 3, height))
output_img.paste(img, (0, 0))
output_img.paste(output, (width, 0))
output_img.paste(target_image, (width * 2, 0))
output_img.save(save_path)
with torch.no_grad():
for i, img_url in enumerate(self.sample_sources):
img = exif_transpose(Image.open(img_url))
img = img.convert('RGB')
# crop if not square
if img.width != img.height:
min_dim = min(img.width, img.height)
img = img.crop((0, 0, min_dim, min_dim))
# resize
img = img.resize((self.resolution * self.zoom, self.resolution * self.zoom), resample=Image.BICUBIC)
target_image = img
# downscale the image input
img = img.resize((self.resolution, self.resolution), resample=Image.BICUBIC)
# downscale the image input
img = IMAGE_TRANSFORMS(img).unsqueeze(0).to(self.device, dtype=self.esrgan_dtype)
img = img
step_num = ''
if step is not None:
# zero-pad 9 digits
step_num = f"_{str(step).zfill(9)}"
seconds_since_epoch = int(time.time())
# zero-pad 2 digits
i_str = str(i).zfill(2)
filename = f"{seconds_since_epoch}{step_num}_{i_str}.jpg"
process_and_save(img, target_image, os.path.join(sample_folder, filename))
if batch is not None:
batch_targets = batch[0].detach()
batch_inputs = batch[1].detach()
batch_targets = torch.chunk(batch_targets, batch_targets.shape[0], dim=0)
batch_inputs = torch.chunk(batch_inputs, batch_inputs.shape[0], dim=0)
for i in range(len(batch_inputs)):
if step is not None:
# zero-pad 9 digits
step_num = f"_{str(step).zfill(9)}"
seconds_since_epoch = int(time.time())
# zero-pad 2 digits
i_str = str(i).zfill(2)
filename = f"{seconds_since_epoch}{step_num}_{i_str}.jpg"
process_and_save(batch_inputs[i], batch_targets[i], os.path.join(batch_sample_folder, filename))
self.model.train()
def load_model(self):
state_dict = None
path_to_load = self.pretrained_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}*.pth"))
if files and len(files) > 0:
latest_file = max(files, key=os.path.getmtime)
print(f" - Latest checkpoint is: {latest_file}")
path_to_load = latest_file
# todo update step and epoch count
elif self.pretrained_path is None:
self.print(f" - No checkpoint found, starting from scratch")
else:
self.print(f" - No checkpoint found, loading pretrained model")
self.print(f" - path: {path_to_load}")
if path_to_load is not None:
self.print(f" - Loading pretrained checkpoint: {path_to_load}")
# if ends with pth then assume pytorch checkpoint
if path_to_load.endswith('.pth') or path_to_load.endswith('.pt'):
state_dict = torch.load(path_to_load, map_location=self.device)
elif path_to_load.endswith('.safetensors'):
state_dict_raw = load_file(path_to_load)
# make ordered dict as most things need it
state_dict = OrderedDict()
for key in esrgan_safetensors_keys:
state_dict[key] = state_dict_raw[key]
else:
raise Exception(f"Unknown file extension for checkpoint: {path_to_load}")
# todo determine architecture from checkpoint
self.model = ESRGAN(
state_dict
).to(self.device, dtype=self.esrgan_dtype)
# set the model to training mode
self.model.train()
self.model.requires_grad_(True)
def run(self):
super().run()
self.load_datasets()
steps_per_step = (self.critic.num_critic_per_gen + 1)
max_step_epochs = self.max_steps // (len(self.data_loader) // steps_per_step)
num_epochs = self.epochs
if num_epochs is None or num_epochs > max_step_epochs:
num_epochs = max_step_epochs
max_epoch_steps = len(self.data_loader) * num_epochs * steps_per_step
num_steps = self.max_steps
if num_steps is None or num_steps > max_epoch_steps:
num_steps = max_epoch_steps
self.max_steps = num_steps
self.epochs = num_epochs
start_step = self.step_num
self.first_step = start_step
self.print(f"Training ESRGAN model:")
self.print(f" - Training folder: {self.training_folder}")
self.print(f" - Batch size: {self.batch_size}")
self.print(f" - Learning rate: {self.learning_rate}")
self.print(f" - Epochs: {num_epochs}")
self.print(f" - Max steps: {self.max_steps}")
# load model
self.load_model()
params = self.model.parameters()
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
self.setup_vgg19()
self.vgg_19.requires_grad_(False)
self.vgg_19.eval()
if self.use_critic:
self.critic.setup()
optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
optimizer_params=self.optimizer_params)
# setup scheduler
# todo allow other schedulers
scheduler = torch.optim.lr_scheduler.ConstantLR(
optimizer,
total_iters=num_steps,
factor=1,
verbose=False
)
# setup tqdm progress bar
self.progress_bar = tqdm(
total=num_steps,
desc='Training ESRGAN',
leave=True
)
blank_losses = OrderedDict({
"total": [],
"style": [],
"content": [],
"mse": [],
"kl": [],
"tv": [],
"ptn": [],
"crD": [],
"crG": [],
})
epoch_losses = copy.deepcopy(blank_losses)
log_losses = copy.deepcopy(blank_losses)
print("Generating baseline samples")
self.sample(step=0)
# range start at self.epoch_num go to self.epochs
critic_losses = []
for epoch in range(self.epoch_num, self.epochs, 1):
if self.step_num >= self.max_steps:
break
flush()
for targets, inputs in self.data_loader:
if self.step_num >= self.max_steps:
break
with torch.no_grad():
is_critic_only_step = False
if self.use_critic and 1 / (self.critic.num_critic_per_gen + 1) < np.random.uniform():
is_critic_only_step = True
targets = targets.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1).detach()
inputs = inputs.to(self.device, dtype=self.esrgan_dtype).clamp(0, 1).detach()
optimizer.zero_grad()
# dont do grads here for critic step
do_grad = not is_critic_only_step
with torch.set_grad_enabled(do_grad):
pred = self.model(inputs)
pred = pred.to(self.device, dtype=self.torch_dtype).clamp(0, 1)
targets = targets.to(self.device, dtype=self.torch_dtype).clamp(0, 1)
if torch.isnan(pred).any():
raise ValueError('pred has nan values')
if torch.isnan(targets).any():
raise ValueError('targets has nan values')
# Run through VGG19
if self.style_weight > 0 or self.content_weight > 0 or self.use_critic:
stacked = torch.cat([pred, targets], dim=0)
# stacked = (stacked / 2 + 0.5).clamp(0, 1)
stacked = stacked.clamp(0, 1)
self.vgg_19(stacked)
# make sure we dont have nans
if torch.isnan(self.vgg19_pool_4.tensor).any():
raise ValueError('vgg19_pool_4 has nan values')
if is_critic_only_step:
critic_d_loss = self.critic.step(self.vgg19_pool_4.tensor.detach())
critic_losses.append(critic_d_loss)
# don't do generator step
continue
else:
# doing a regular step
if len(critic_losses) == 0:
critic_d_loss = 0
else:
critic_d_loss = sum(critic_losses) / len(critic_losses)
style_loss = self.get_style_loss() * self.style_weight
content_loss = self.get_content_loss() * self.content_weight
mse_loss = self.get_mse_loss(pred, targets) * self.mse_weight
tv_loss = self.get_tv_loss(pred, targets) * self.tv_weight
pattern_loss = self.get_pattern_loss(pred, targets) * self.pattern_weight
if self.use_critic:
critic_gen_loss = self.critic.get_critic_loss(self.vgg19_pool_4.tensor) * self.critic_weight
else:
critic_gen_loss = torch.tensor(0.0, device=self.device, dtype=self.torch_dtype)
loss = style_loss + content_loss + mse_loss + tv_loss + critic_gen_loss + pattern_loss
# make sure non nan
if torch.isnan(loss):
raise ValueError('loss is nan')
# Backward pass and optimization
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
optimizer.step()
scheduler.step()
# update progress bar
loss_value = loss.item()
# get exponent like 3.54e-4
loss_string = f"loss: {loss_value:.2e}"
if self.content_weight > 0:
loss_string += f" cnt: {content_loss.item():.2e}"
if self.style_weight > 0:
loss_string += f" sty: {style_loss.item():.2e}"
if self.mse_weight > 0:
loss_string += f" mse: {mse_loss.item():.2e}"
if self.tv_weight > 0:
loss_string += f" tv: {tv_loss.item():.2e}"
if self.pattern_weight > 0:
loss_string += f" ptn: {pattern_loss.item():.2e}"
if self.use_critic and self.critic_weight > 0:
loss_string += f" crG: {critic_gen_loss.item():.2e}"
if self.use_critic:
loss_string += f" crD: {critic_d_loss:.2e}"
if self.optimizer_type.startswith('dadaptation') or self.optimizer_type.startswith('prodigy'):
learning_rate = (
optimizer.param_groups[0]["d"] *
optimizer.param_groups[0]["lr"]
)
else:
learning_rate = optimizer.param_groups[0]['lr']
lr_critic_string = ''
if self.use_critic:
lr_critic = self.critic.get_lr()
lr_critic_string = f" lrC: {lr_critic:.1e}"
self.progress_bar.set_postfix_str(f"lr: {learning_rate:.1e}{lr_critic_string} {loss_string}")
self.progress_bar.set_description(f"E: {epoch}")
self.progress_bar.update(1)
epoch_losses["total"].append(loss_value)
epoch_losses["style"].append(style_loss.item())
epoch_losses["content"].append(content_loss.item())
epoch_losses["mse"].append(mse_loss.item())
epoch_losses["tv"].append(tv_loss.item())
epoch_losses["ptn"].append(pattern_loss.item())
epoch_losses["crG"].append(critic_gen_loss.item())
epoch_losses["crD"].append(critic_d_loss)
log_losses["total"].append(loss_value)
log_losses["style"].append(style_loss.item())
log_losses["content"].append(content_loss.item())
log_losses["mse"].append(mse_loss.item())
log_losses["tv"].append(tv_loss.item())
log_losses["ptn"].append(pattern_loss.item())
log_losses["crG"].append(critic_gen_loss.item())
log_losses["crD"].append(critic_d_loss)
# don't do on first step
if self.step_num != start_step:
if self.sample_every and self.step_num % self.sample_every == 0:
# print above the progress bar
self.print(f"Sampling at step {self.step_num}")
self.sample(self.step_num, batch=[targets, inputs])
if self.save_every and self.step_num % self.save_every == 0:
# print above the progress bar
self.print(f"Saving at step {self.step_num}")
self.save(self.step_num)
if self.log_every and self.step_num % self.log_every == 0:
# log to tensorboard
if self.writer is not None:
# get avg loss
for key in log_losses:
log_losses[key] = sum(log_losses[key]) / (len(log_losses[key]) + 1e-6)
# if log_losses[key] > 0:
self.writer.add_scalar(f"loss/{key}", log_losses[key], self.step_num)
# reset log losses
log_losses = copy.deepcopy(blank_losses)
self.step_num += 1
# end epoch
if self.writer is not None:
eps = 1e-6
# get avg loss
for key in epoch_losses:
epoch_losses[key] = sum(log_losses[key]) / (len(log_losses[key]) + eps)
if epoch_losses[key] > 0:
self.writer.add_scalar(f"epoch loss/{key}", epoch_losses[key], epoch)
# reset epoch losses
epoch_losses = copy.deepcopy(blank_losses)
self.save()

View File

@@ -1,76 +0,0 @@
# ref:
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
import time
from collections import OrderedDict
import os
from toolkit.config_modules import SliderConfig
from toolkit.paths import REPOS_ROOT
import sys
sys.path.append(REPOS_ROOT)
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
import gc
import torch
from leco import train_util, model_util
from leco.prompt_util import PromptEmbedsCache
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
def flush():
torch.cuda.empty_cache()
gc.collect()
class LoRAHack:
def __init__(self, **kwargs):
self.type = kwargs.get('type', 'suppression')
class TrainLoRAHack(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.hack_config = LoRAHack(**self.get_conf('hack', {}))
def hook_before_train_loop(self):
# we don't need text encoder so move it to cpu
self.sd.text_encoder.to("cpu")
flush()
# end hook_before_train_loop
if self.hack_config.type == 'suppression':
# set all params to self.current_suppression
params = self.network.parameters()
for param in params:
# get random noise for each param
noise = torch.randn_like(param) - 0.5
# apply noise to param
param.data = noise * 0.001
def supress_loop(self):
dtype = get_torch_dtype(self.train_config.dtype)
loss_dict = OrderedDict(
{'sup': 0.0}
)
# increase noise
for param in self.network.parameters():
# get random noise for each param
noise = torch.randn_like(param) - 0.5
# apply noise to param
param.data = param.data + noise * 0.001
return loss_dict
def hook_train_loop(self):
if self.hack_config.type == 'suppression':
return self.supress_loop()
else:
raise NotImplementedError(f'unknown hack type: {self.hack_config.type}')
# end hook_train_loop

View File

@@ -1,22 +1,14 @@
# ref:
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
import time
from collections import OrderedDict
import glob
import os
from typing import Optional
from collections import OrderedDict
import random
from typing import Optional, List
from safetensors.torch import load_file, save_file
from safetensors.torch import save_file, load_file
from tqdm import tqdm
from toolkit.config_modules import SliderConfig
from toolkit.layers import ReductionKernel
from toolkit.paths import REPOS_ROOT
import sys
from toolkit.stable_diffusion_model import PromptEmbeds
sys.path.append(REPOS_ROOT)
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
import gc
from toolkit import train_tools
@@ -38,12 +30,10 @@ class RescaleConfig:
):
self.from_resolution = kwargs.get('from_resolution', 512)
self.scale = kwargs.get('scale', 0.5)
self.prompt_file = kwargs.get('prompt_file', None)
self.prompt_tensors = kwargs.get('prompt_tensors', None)
self.latent_tensor_dir = kwargs.get('latent_tensor_dir', None)
self.num_latent_tensors = kwargs.get('num_latent_tensors', 1000)
self.to_resolution = kwargs.get('to_resolution', int(self.from_resolution * self.scale))
if self.prompt_file is None:
raise ValueError("prompt_file is required")
self.prompt_dropout = kwargs.get('prompt_dropout', 0.1)
class PromptEmbedsCache:
@@ -61,12 +51,12 @@ class PromptEmbedsCache:
class TrainSDRescaleProcess(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
# pass our custom pipeline to super so it sets it up
super().__init__(process_id, job, config)
self.step_num = 0
self.start_step = 0
self.device = self.get_conf('device', self.job.device)
self.device_torch = torch.device(self.device)
self.prompt_cache = PromptEmbedsCache()
self.rescale_config = RescaleConfig(**self.get_conf('rescale', required=True))
self.reduce_size_fn = ReductionKernel(
in_channels=4,
@@ -74,202 +64,211 @@ class TrainSDRescaleProcess(BaseSDTrainProcess):
dtype=get_torch_dtype(self.train_config.dtype),
device=self.device_torch,
)
self.prompt_txt_list = []
self.latent_paths: List[str] = []
self.empty_embedding: PromptEmbeds = None
def before_model_load(self):
pass
def get_latent_tensors(self):
dtype = get_torch_dtype(self.train_config.dtype)
num_to_generate = 0
# check if dir exists
if not os.path.exists(self.rescale_config.latent_tensor_dir):
os.makedirs(self.rescale_config.latent_tensor_dir)
num_to_generate = self.rescale_config.num_latent_tensors
else:
# find existing
current_tensor_list = glob.glob(os.path.join(self.rescale_config.latent_tensor_dir, "*.safetensors"))
num_to_generate = self.rescale_config.num_latent_tensors - len(current_tensor_list)
self.latent_paths = current_tensor_list
if num_to_generate > 0:
print(f"Generating {num_to_generate}/{self.rescale_config.num_latent_tensors} latent tensors")
# unload other model
self.sd.unet.to('cpu')
# load aux network
self.sd_parent = StableDiffusion(
self.device_torch,
model_config=self.model_config,
dtype=self.train_config.dtype,
)
self.sd_parent.load_model()
self.sd_parent.unet.to(self.device_torch, dtype=dtype)
# we dont need text encoder for this
del self.sd_parent.text_encoder
del self.sd_parent.tokenizer
self.sd_parent.unet.eval()
self.sd_parent.unet.requires_grad_(False)
# save current seed state for training
rng_state = torch.get_rng_state()
cuda_rng_state = torch.cuda.get_rng_state() if torch.cuda.is_available() else None
text_embeddings = train_tools.concat_prompt_embeddings(
self.empty_embedding, # unconditional (negative prompt)
self.empty_embedding, # conditional (positive prompt)
self.train_config.batch_size,
)
torch.set_default_device(self.device_torch)
for i in tqdm(range(num_to_generate)):
dtype = get_torch_dtype(self.train_config.dtype)
# get a random seed
seed = torch.randint(0, 2 ** 32, (1,)).item()
# zero pad seed string to max length
seed_string = str(seed).zfill(10)
# set seed
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
# # ger a random number of steps
timesteps_to = self.train_config.max_denoising_steps
# set the scheduler to the number of steps
self.sd.noise_scheduler.set_timesteps(
timesteps_to, device=self.device_torch
)
noise = self.sd.get_latent_noise(
pixel_height=self.rescale_config.from_resolution,
pixel_width=self.rescale_config.from_resolution,
batch_size=self.train_config.batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
# get latents
latents = noise * self.sd.noise_scheduler.init_noise_sigma
latents = latents.to(self.device_torch, dtype=dtype)
# get random guidance scale from 1.0 to 10.0 (CFG)
guidance_scale = torch.rand(1).item() * 9.0 + 1.0
# do a timestep of 1
timestep = 1
noise_pred_target = self.sd_parent.predict_noise(
latents,
text_embeddings=text_embeddings,
timestep=timestep,
guidance_scale=guidance_scale
)
# build state dict
state_dict = OrderedDict()
state_dict['noise_pred_target'] = noise_pred_target.to('cpu', dtype=torch.float16)
state_dict['latents'] = latents.to('cpu', dtype=torch.float16)
state_dict['guidance_scale'] = torch.tensor(guidance_scale).to('cpu', dtype=torch.float16)
state_dict['timestep'] = torch.tensor(timestep).to('cpu', dtype=torch.float16)
state_dict['timesteps_to'] = torch.tensor(timesteps_to).to('cpu', dtype=torch.float16)
state_dict['seed'] = torch.tensor(seed).to('cpu', dtype=torch.float32) # must be float 32 to prevent overflow
file_name = f"{seed_string}_{i}.safetensors"
file_path = os.path.join(self.rescale_config.latent_tensor_dir, file_name)
save_file(state_dict, file_path)
self.latent_paths.append(file_path)
print("Removing parent model")
# delete parent
del self.sd_parent
flush()
torch.set_rng_state(rng_state)
if cuda_rng_state is not None:
torch.cuda.set_rng_state(cuda_rng_state)
self.sd.unet.to(self.device_torch, dtype=dtype)
def hook_before_train_loop(self):
self.print(f"Loading prompt file from {self.rescale_config.prompt_file}")
# encode our empty prompt
self.empty_embedding = self.sd.encode_prompt("")
self.empty_embedding = self.empty_embedding.to(self.device_torch,
dtype=get_torch_dtype(self.train_config.dtype))
# read line by line from file
with open(self.rescale_config.prompt_file, 'r') as f:
self.prompt_txt_list = f.readlines()
# clean empty lines
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
self.print(f"Loaded {len(self.prompt_txt_list)} prompts. Encoding them..")
cache = PromptEmbedsCache()
# get encoded latents for our prompts
with torch.no_grad():
if self.rescale_config.prompt_tensors is not None:
# check to see if it exists
if os.path.exists(self.rescale_config.prompt_tensors):
# load it.
self.print(f"Loading prompt tensors from {self.rescale_config.prompt_tensors}")
prompt_tensors = load_file(self.rescale_config.prompt_tensors, device='cpu')
# add them to the cache
for prompt_txt, prompt_tensor in prompt_tensors.items():
if prompt_txt.startswith("te:"):
prompt = prompt_txt[3:]
# text_embeds
text_embeds = prompt_tensor
pooled_embeds = None
# find pool embeds
if f"pe:{prompt}" in prompt_tensors:
pooled_embeds = prompt_tensors[f"pe:{prompt}"]
# make it
prompt_embeds = PromptEmbeds([text_embeds, pooled_embeds])
cache[prompt] = prompt_embeds.to(device='cpu', dtype=torch.float32)
if len(cache.prompts) == 0:
print("Prompt tensors not found. Encoding prompts..")
neutral = ""
# encode neutral
cache[neutral] = self.sd.encode_prompt(neutral)
for prompt in tqdm(self.prompt_txt_list, desc="Encoding prompts", leave=False):
# build the cache
if cache[prompt] is None:
cache[prompt] = self.sd.encode_prompt(prompt).to(device="cpu", dtype=torch.float32)
if self.rescale_config.prompt_tensors:
print(f"Saving prompt tensors to {self.rescale_config.prompt_tensors}")
state_dict = {}
for prompt_txt, prompt_embeds in cache.prompts.items():
state_dict[f"te:{prompt_txt}"] = prompt_embeds.text_embeds.to("cpu", dtype=get_torch_dtype('fp16'))
if prompt_embeds.pooled_embeds is not None:
state_dict[f"pe:{prompt_txt}"] = prompt_embeds.pooled_embeds.to("cpu", dtype=get_torch_dtype('fp16'))
save_file(state_dict, self.rescale_config.prompt_tensors)
self.print("Encoding complete.")
# move to cpu to save vram
# We don't need text encoder anymore, but keep it on cpu for sampling
# if text encoder is list
# Move train model encoder to cpu
if isinstance(self.sd.text_encoder, list):
for encoder in self.sd.text_encoder:
encoder.to("cpu")
encoder.to('cpu')
encoder.eval()
encoder.requires_grad_(False)
else:
self.sd.text_encoder.to("cpu")
self.prompt_cache = cache
self.sd.text_encoder.to('cpu')
self.sd.text_encoder.eval()
self.sd.text_encoder.requires_grad_(False)
# self.sd.unet.to('cpu')
flush()
self.get_latent_tensors()
flush()
# end hook_before_train_loop
def hook_train_loop(self):
def hook_train_loop(self, batch):
dtype = get_torch_dtype(self.train_config.dtype)
# get random encoded prompt from cache
prompt_txt = self.prompt_txt_list[
torch.randint(0, len(self.prompt_txt_list), (1,)).item()
]
prompt = self.prompt_cache[prompt_txt].to(device=self.device_torch, dtype=dtype)
neutral = self.prompt_cache[""].to(device=self.device_torch, dtype=dtype)
if prompt is None:
raise ValueError(f"Prompt {prompt_txt} is not in cache")
prompt_batch = train_tools.concat_prompt_embeddings(
prompt,
neutral,
self.train_config.batch_size,
)
noise_scheduler = self.sd.noise_scheduler
optimizer = self.optimizer
lr_scheduler = self.lr_scheduler
loss_function = torch.nn.MSELoss()
def get_noise_pred(p, n, gs, cts, dn):
return self.predict_noise(
latents=dn,
text_embeddings=train_tools.concat_prompt_embeddings(
p, # unconditional
n, # positive
self.train_config.batch_size,
),
timestep=cts,
guidance_scale=gs,
)
# train it
# Begin gradient accumulation
self.sd.unet.train()
self.sd.unet.requires_grad_(True)
self.sd.unet.to(self.device_torch, dtype=dtype)
with torch.no_grad():
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
self.optimizer.zero_grad()
# # ger a random number of steps
timesteps_to = torch.randint(
1, self.train_config.max_denoising_steps, (1,)
).item()
# pick random latent tensor
latent_path = random.choice(self.latent_paths)
latent_tensor = load_file(latent_path)
# get noise
noise = self.get_latent_noise(
pixel_height=self.rescale_config.from_resolution,
pixel_width=self.rescale_config.from_resolution,
).to(self.device_torch, dtype=dtype)
noise_pred_target = (latent_tensor['noise_pred_target']).to(self.device_torch, dtype=dtype)
latents = (latent_tensor['latents']).to(self.device_torch, dtype=dtype)
guidance_scale = (latent_tensor['guidance_scale']).item()
timestep = int((latent_tensor['timestep']).item())
timesteps_to = int((latent_tensor['timesteps_to']).item())
# seed = int((latent_tensor['seed']).item())
# get latents
latents = noise * self.sd.noise_scheduler.init_noise_sigma
latents = latents.to(self.device_torch, dtype=dtype)
#
# # predict without network
# assert self.network.is_active is False
# denoised_latents = self.diffuse_some_steps(
# latents, # pass simple noise latents
# prompt_batch,
# start_timesteps=0,
# total_timesteps=timesteps_to,
# guidance_scale=3,
# )
# noise_scheduler.set_timesteps(1000)
#
# current_timestep = noise_scheduler.timesteps[
# int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
# ]
current_timestep = 0
denoised_latents = latents
# get noise prediction at full scale
from_prediction = get_noise_pred(
prompt, neutral, 1, current_timestep, denoised_latents
text_embeddings = train_tools.concat_prompt_embeddings(
self.empty_embedding, # unconditional (negative prompt)
self.empty_embedding, # conditional (positive prompt)
self.train_config.batch_size,
)
self.sd.noise_scheduler.set_timesteps(
timesteps_to, device=self.device_torch
)
reduced_from_prediction = self.reduce_size_fn(from_prediction).to("cpu", dtype=torch.float32)
denoised_target = self.sd.noise_scheduler.step(noise_pred_target, timestep, latents).prev_sample
# get noise prediction at reduced scale
to_denoised_latents = self.reduce_size_fn(denoised_latents)
# get the reduced latents
# reduced_pred = self.reduce_size_fn(noise_pred_target.detach())
denoised_target = self.reduce_size_fn(denoised_target.detach())
reduced_latents = self.reduce_size_fn(latents.detach())
# start gradient
optimizer.zero_grad()
self.network.multiplier = 1.0
with self.network:
assert self.network.is_active is True
to_prediction = get_noise_pred(
prompt, neutral, 1, current_timestep, to_denoised_latents
).to("cpu", dtype=torch.float32)
reduced_from_prediction.requires_grad = False
from_prediction.requires_grad = False
loss = loss_function(
reduced_from_prediction,
to_prediction,
denoised_target.requires_grad = False
self.optimizer.zero_grad()
noise_pred_train = self.sd.predict_noise(
reduced_latents,
text_embeddings=text_embeddings,
timestep=timestep,
guidance_scale=guidance_scale
)
denoised_pred = self.sd.noise_scheduler.step(noise_pred_train, timestep, reduced_latents).prev_sample
loss = loss_function(denoised_pred, denoised_target)
loss_float = loss.item()
loss = loss.to(self.device_torch)
loss.backward()
optimizer.step()
lr_scheduler.step()
self.optimizer.step()
self.lr_scheduler.step()
self.optimizer.zero_grad()
del (
reduced_from_prediction,
from_prediction,
to_denoised_latents,
to_prediction,
latents,
)
flush()
# reset network
self.network.multiplier = 1.0
loss_dict = OrderedDict(
{'loss': loss_float},
)

View File

@@ -1,30 +1,29 @@
# ref:
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
import time
from collections import OrderedDict
import copy
import os
from typing import Optional
import random
from collections import OrderedDict
from typing import Union
from PIL import Image
from diffusers import T2IAdapter
from torchvision.transforms import transforms
from tqdm import tqdm
from toolkit.basic import value_map
from toolkit.config_modules import SliderConfig
from toolkit.paths import REPOS_ROOT
import sys
from toolkit.stable_diffusion_model import PromptEmbeds
sys.path.append(REPOS_ROOT)
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.sd_device_states_presets import get_train_sd_device_state_preset
from toolkit.train_tools import get_torch_dtype, apply_snr_weight, apply_learnable_snr_gos
import gc
from toolkit import train_tools
from toolkit.prompt_utils import \
EncodedPromptPair, ACTION_TYPES_SLIDER, \
EncodedAnchor, concat_prompt_pairs, \
concat_anchors, PromptEmbedsCache, encode_prompts_to_cache, build_prompt_pair_batch_from_cache, split_anchors, \
split_prompt_pairs
import torch
from leco import train_util, model_util
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
class ACTION_TYPES_SLIDER:
ERASE_NEGATIVE = 0
ENHANCE_NEGATIVE = 1
from .BaseSDTrainProcess import BaseSDTrainProcess
def flush():
@@ -32,58 +31,15 @@ def flush():
gc.collect()
class EncodedPromptPair:
def __init__(
self,
target_class,
positive,
negative,
neutral,
width=512,
height=512,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=1.0,
weight=1.0
):
self.target_class = target_class
self.positive = positive
self.negative = negative
self.neutral = neutral
self.width = width
self.height = height
self.action: int = action
self.multiplier = multiplier
self.weight = weight
class PromptEmbedsCache: # 使いまわしたいので
prompts: dict[str, PromptEmbeds] = {}
def __setitem__(self, __name: str, __value: PromptEmbeds) -> None:
self.prompts[__name] = __value
def __getitem__(self, __name: str) -> Optional[PromptEmbeds]:
if __name in self.prompts:
return self.prompts[__name]
else:
return None
class EncodedAnchor:
def __init__(
self,
prompt,
neg_prompt,
multiplier=1.0
):
self.prompt = prompt
self.neg_prompt = neg_prompt
self.multiplier = multiplier
adapter_transforms = transforms.Compose([
transforms.ToTensor(),
])
class TrainSliderProcess(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.prompt_txt_list = None
self.step_num = 0
self.start_step = 0
self.device = self.get_conf('device', self.job.device)
@@ -92,102 +48,119 @@ class TrainSliderProcess(BaseSDTrainProcess):
self.prompt_cache = PromptEmbedsCache()
self.prompt_pairs: list[EncodedPromptPair] = []
self.anchor_pairs: list[EncodedAnchor] = []
# keep track of prompt chunk size
self.prompt_chunk_size = 1
# check if we have more targets than steps
# this can happen because of permutation son shuffling
if len(self.slider_config.targets) > self.train_config.steps:
# trim targets
self.slider_config.targets = self.slider_config.targets[:self.train_config.steps]
# get presets
self.eval_slider_device_state = get_train_sd_device_state_preset(
self.device_torch,
train_unet=False,
train_text_encoder=False,
cached_latents=self.is_latents_cached,
train_lora=False,
train_adapter=False,
train_embedding=False,
)
self.train_slider_device_state = get_train_sd_device_state_preset(
self.device_torch,
train_unet=self.train_config.train_unet,
train_text_encoder=False,
cached_latents=self.is_latents_cached,
train_lora=True,
train_adapter=False,
train_embedding=False,
)
def before_model_load(self):
pass
def hook_before_train_loop(self):
# read line by line from file
if self.slider_config.prompt_file:
self.print(f"Loading prompt file from {self.slider_config.prompt_file}")
with open(self.slider_config.prompt_file, 'r', encoding='utf-8') as f:
self.prompt_txt_list = f.readlines()
# clean empty lines
self.prompt_txt_list = [line.strip() for line in self.prompt_txt_list if len(line.strip()) > 0]
self.print(f"Found {len(self.prompt_txt_list)} prompts.")
if not self.slider_config.prompt_tensors:
print(f"Prompt tensors not found. Building prompt tensors for {self.train_config.steps} steps.")
# shuffle
random.shuffle(self.prompt_txt_list)
# trim to max steps
self.prompt_txt_list = self.prompt_txt_list[:self.train_config.steps]
# trim list to our max steps
cache = PromptEmbedsCache()
prompt_pairs: list[EncodedPromptPair] = []
print(f"Building prompt cache")
# get encoded latents for our prompts
with torch.no_grad():
neutral = ""
for target in self.slider_config.targets:
# build the cache
for prompt in [
target.target_class,
target.positive,
target.negative,
neutral # empty neutral
]:
if cache[prompt] is None:
cache[prompt] = self.sd.encode_prompt(prompt)
for resolution in self.slider_config.resolutions:
width, height = resolution
erase_negative = len(target.positive.strip()) == 0
enhance_positive = len(target.negative.strip()) == 0
# list of neutrals. Can come from file or be empty
neutral_list = self.prompt_txt_list if self.prompt_txt_list is not None else [""]
both = not erase_negative and not enhance_positive
# build the prompts to cache
prompts_to_cache = []
for neutral in neutral_list:
for target in self.slider_config.targets:
prompt_list = [
f"{target.target_class}", # target_class
f"{target.target_class} {neutral}", # target_class with neutral
f"{target.positive}", # positive_target
f"{target.positive} {neutral}", # positive_target with neutral
f"{target.negative}", # negative_target
f"{target.negative} {neutral}", # negative_target with neutral
f"{neutral}", # neutral
f"{target.positive} {target.negative}", # both targets
f"{target.negative} {target.positive}", # both targets reverse
]
prompts_to_cache += prompt_list
if erase_negative and enhance_positive:
raise ValueError("target must have at least one of positive or negative or both")
# for slider we need to have an enhancer, an eraser, and then
# an inverse with negative weights to balance the network
# if we don't do this, we will get different contrast and focus.
# we only perform actions of enhancing and erasing on the negative
# todo work on way to do all of this in one shot
# remove duplicates
prompts_to_cache = list(dict.fromkeys(prompts_to_cache))
if both or erase_negative:
prompt_pairs += [
# erase standard
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.positive],
negative=cache[target.negative],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=target.multiplier,
weight=target.weight
),
]
if both or enhance_positive:
prompt_pairs += [
# enhance standard, swap pos neg
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.negative],
negative=cache[target.positive],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
multiplier=target.multiplier,
weight=target.weight
),
]
if both or enhance_positive:
prompt_pairs += [
# erase inverted
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.negative],
negative=cache[target.positive],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=target.multiplier * -1.0,
weight=target.weight
),
]
if both or erase_negative:
prompt_pairs += [
# enhance inverted
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.positive],
negative=cache[target.negative],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
multiplier=target.multiplier * -1.0,
weight=target.weight
),
]
# trim to max steps if max steps is lower than prompt count
# todo, this can break if we have more targets than steps, should be fixed, by reducing permuations, but could stil happen with low steps
# prompts_to_cache = prompts_to_cache[:self.train_config.steps]
# encode them
cache = encode_prompts_to_cache(
prompt_list=prompts_to_cache,
sd=self.sd,
cache=cache,
prompt_tensor_file=self.slider_config.prompt_tensors
)
prompt_pairs = []
prompt_batches = []
for neutral in tqdm(neutral_list, desc="Building Prompt Pairs", leave=False):
for target in self.slider_config.targets:
prompt_pair_batch = build_prompt_pair_batch_from_cache(
cache=cache,
target=target,
neutral=neutral,
)
if self.slider_config.batch_full_slide:
# concat the prompt pairs
# this allows us to run the entire 4 part process in one shot (for slider)
self.prompt_chunk_size = 4
concat_prompt_pair_batch = concat_prompt_pairs(prompt_pair_batch).to('cpu')
prompt_pairs += [concat_prompt_pair_batch]
else:
self.prompt_chunk_size = 1
# do them one at a time (probably not necessary after new optimizations)
prompt_pairs += [x.to('cpu') for x in prompt_pair_batch]
# setup anchors
anchor_pairs = []
@@ -200,13 +173,26 @@ class TrainSliderProcess(BaseSDTrainProcess):
if cache[prompt] == None:
cache[prompt] = self.sd.encode_prompt(prompt)
anchor_batch = []
# we get the prompt pair multiplier from first prompt pair
# since they are all the same. We need to match their network polarity
prompt_pair_multipliers = prompt_pairs[0].multiplier_list
for prompt_multiplier in prompt_pair_multipliers:
# match the network multiplier polarity
anchor_scalar = 1.0 if prompt_multiplier > 0 else -1.0
anchor_batch += [
EncodedAnchor(
prompt=cache[anchor.prompt],
neg_prompt=cache[anchor.neg_prompt],
multiplier=anchor.multiplier * anchor_scalar
)
]
anchor_pairs += [
EncodedAnchor(
prompt=cache[anchor.prompt],
neg_prompt=cache[anchor.neg_prompt],
multiplier=anchor.multiplier
)
concat_anchors(anchor_batch).to('cpu')
]
if len(anchor_pairs) > 0:
self.anchor_pairs = anchor_pairs
# move to cpu to save vram
# We don't need text encoder anymore, but keep it on cpu for sampling
@@ -218,78 +204,198 @@ class TrainSliderProcess(BaseSDTrainProcess):
self.sd.text_encoder.to("cpu")
self.prompt_cache = cache
self.prompt_pairs = prompt_pairs
self.anchor_pairs = anchor_pairs
# self.anchor_pairs = anchor_pairs
flush()
if self.data_loader is not None:
# we will have images, prep the vae
self.sd.vae.eval()
self.sd.vae.to(self.device_torch)
# end hook_before_train_loop
def hook_train_loop(self):
dtype = get_torch_dtype(self.train_config.dtype)
def before_dataset_load(self):
if self.slider_config.use_adapter == 'depth':
print(f"Loading T2I Adapter for depth")
# called before LoRA network is loaded but after model is loaded
# attach the adapter here so it is there before we load the network
adapter_path = 'TencentARC/t2iadapter_depth_sd15v2'
if self.model_config.is_xl:
adapter_path = 'TencentARC/t2i-adapter-depth-midas-sdxl-1.0'
# get a random pair
prompt_pair: EncodedPromptPair = self.prompt_pairs[
torch.randint(0, len(self.prompt_pairs), (1,)).item()
]
print(f"Loading T2I Adapter from {adapter_path}")
height = prompt_pair.height
width = prompt_pair.width
target_class = prompt_pair.target_class
neutral = prompt_pair.neutral
negative = prompt_pair.negative
positive = prompt_pair.positive
weight = prompt_pair.weight
multiplier = prompt_pair.multiplier
# dont name this adapter since we are not training it
self.t2i_adapter = T2IAdapter.from_pretrained(
adapter_path, torch_dtype=get_torch_dtype(self.train_config.dtype), varient="fp16"
).to(self.device_torch)
self.t2i_adapter.eval()
self.t2i_adapter.requires_grad_(False)
flush()
@torch.no_grad()
def get_adapter_images(self, batch: Union[None, 'DataLoaderBatchDTO']):
img_ext_list = ['.jpg', '.jpeg', '.png', '.webp']
adapter_folder_path = self.slider_config.adapter_img_dir
adapter_images = []
# loop through images
for file_item in batch.file_items:
img_path = file_item.path
file_name_no_ext = os.path.basename(img_path).split('.')[0]
# find the image
for ext in img_ext_list:
if os.path.exists(os.path.join(adapter_folder_path, file_name_no_ext + ext)):
adapter_images.append(os.path.join(adapter_folder_path, file_name_no_ext + ext))
break
width, height = batch.file_items[0].crop_width, batch.file_items[0].crop_height
adapter_tensors = []
# load images with torch transforms
for idx, adapter_image in enumerate(adapter_images):
# we need to centrally crop the largest dimension of the image to match the batch shape after scaling
# to the smallest dimension
img: Image.Image = Image.open(adapter_image)
if img.width > img.height:
# scale down so height is the same as batch
new_height = height
new_width = int(img.width * (height / img.height))
else:
new_width = width
new_height = int(img.height * (width / img.width))
img = img.resize((new_width, new_height))
crop_fn = transforms.CenterCrop((height, width))
# crop the center to match batch
img = crop_fn(img)
img = adapter_transforms(img)
adapter_tensors.append(img)
# stack them
adapter_tensors = torch.stack(adapter_tensors).to(
self.device_torch, dtype=get_torch_dtype(self.train_config.dtype)
)
return adapter_tensors
def hook_train_loop(self, batch: Union['DataLoaderBatchDTO', None]):
# set to eval mode
self.sd.set_device_state(self.eval_slider_device_state)
with torch.no_grad():
dtype = get_torch_dtype(self.train_config.dtype)
# get a random pair
prompt_pair: EncodedPromptPair = self.prompt_pairs[
torch.randint(0, len(self.prompt_pairs), (1,)).item()
]
# move to device and dtype
prompt_pair.to(self.device_torch, dtype=dtype)
# get a random resolution
height, width = self.slider_config.resolutions[
torch.randint(0, len(self.slider_config.resolutions), (1,)).item()
]
if self.train_config.gradient_checkpointing:
# may get disabled elsewhere
self.sd.unet.enable_gradient_checkpointing()
unet = self.sd.unet
noise_scheduler = self.sd.noise_scheduler
optimizer = self.optimizer
lr_scheduler = self.lr_scheduler
loss_function = torch.nn.MSELoss()
def get_noise_pred(p, n, gs, cts, dn):
return self.predict_noise(
pred_kwargs = {}
def get_noise_pred(neg, pos, gs, cts, dn):
down_kwargs = copy.deepcopy(pred_kwargs)
if 'down_block_additional_residuals' in down_kwargs:
dbr_batch_size = down_kwargs['down_block_additional_residuals'][0].shape[0]
if dbr_batch_size != dn.shape[0]:
amount_to_add = int(dn.shape[0] * 2 / dbr_batch_size)
down_kwargs['down_block_additional_residuals'] = [
torch.cat([sample.clone()] * amount_to_add) for sample in
down_kwargs['down_block_additional_residuals']
]
return self.sd.predict_noise(
latents=dn,
text_embeddings=train_tools.concat_prompt_embeddings(
p, # unconditional
n, # positive
neg, # negative prompt
pos, # positive prompt
self.train_config.batch_size,
),
timestep=cts,
guidance_scale=gs,
**down_kwargs
)
# set network multiplier
self.network.multiplier = multiplier
with torch.no_grad():
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
adapter_images = None
self.sd.unet.eval()
self.optimizer.zero_grad()
# for a complete slider, the batch size is 4 to begin with now
true_batch_size = prompt_pair.target_class.text_embeds.shape[0] * self.train_config.batch_size
from_batch = False
if batch is not None:
# traing from a batch of images, not generating ourselves
from_batch = True
noisy_latents, noise, timesteps, conditioned_prompts, imgs = self.process_general_training_batch(batch)
if self.slider_config.adapter_img_dir is not None:
adapter_images = self.get_adapter_images(batch)
adapter_strength_min = 0.9
adapter_strength_max = 1.0
# ger a random number of steps
timesteps_to = torch.randint(
1, self.train_config.max_denoising_steps, (1,)
).item()
def rand_strength(sample):
adapter_conditioning_scale = torch.rand(
(1,), device=self.device_torch, dtype=dtype
)
# get noise
noise = self.get_latent_noise(
pixel_height=height,
pixel_width=width,
).to(self.device_torch, dtype=dtype)
adapter_conditioning_scale = value_map(
adapter_conditioning_scale,
0.0,
1.0,
adapter_strength_min,
adapter_strength_max
)
return sample.to(self.device_torch, dtype=dtype).detach() * adapter_conditioning_scale
# get latents
latents = noise * self.sd.noise_scheduler.init_noise_sigma
latents = latents.to(self.device_torch, dtype=dtype)
down_block_additional_residuals = self.t2i_adapter(adapter_images)
down_block_additional_residuals = [
rand_strength(sample) for sample in down_block_additional_residuals
]
pred_kwargs['down_block_additional_residuals'] = down_block_additional_residuals
with self.network:
assert self.network.is_active
self.network.multiplier = multiplier
denoised_latents = self.diffuse_some_steps(
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
)
# ger a random number of steps
timesteps_to = torch.randint(
1, self.train_config.max_denoising_steps - 1, (1,)
).item()
# get noise
noise = self.sd.get_latent_noise(
pixel_height=height,
pixel_width=width,
batch_size=true_batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
# get latents
latents = noise * self.sd.noise_scheduler.init_noise_sigma
latents = latents.to(self.device_torch, dtype=dtype)
assert not self.network.is_active
self.sd.unet.eval()
# pass the multiplier list to the network
# 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(
positive, # unconditional
target_class, # target
prompt_pair.positive_target, # unconditional
prompt_pair.target_class, # target
self.train_config.batch_size,
),
start_timesteps=0,
@@ -297,101 +403,282 @@ class TrainSliderProcess(BaseSDTrainProcess):
guidance_scale=3,
)
noise_scheduler.set_timesteps(1000)
current_timestep = noise_scheduler.timesteps[
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
]
noise_scheduler.set_timesteps(1000)
current_timestep_index = int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
current_timestep = noise_scheduler.timesteps[current_timestep_index]
# split the latents into out prompt pair chunks
denoised_latent_chunks = torch.chunk(denoised_latents, self.prompt_chunk_size, dim=0)
denoised_latent_chunks = [x.detach() for x in denoised_latent_chunks]
# flush() # 4.2GB to 3GB on 512x512
mask_multiplier = torch.ones((denoised_latents.shape[0], 1, 1, 1), device=self.device_torch, dtype=dtype)
has_mask = False
if batch and batch.mask_tensor is not None:
with self.timer('get_mask_multiplier'):
# upsampling no supported for bfloat16
mask_multiplier = batch.mask_tensor.to(self.device_torch, dtype=torch.float16).detach()
# scale down to the size of the latents, mask multiplier shape(bs, 1, width, height), noisy_latents shape(bs, channels, width, height)
mask_multiplier = torch.nn.functional.interpolate(
mask_multiplier, size=(noisy_latents.shape[2], noisy_latents.shape[3])
)
# expand to match latents
mask_multiplier = mask_multiplier.expand(-1, noisy_latents.shape[1], -1, -1)
mask_multiplier = mask_multiplier.to(self.device_torch, dtype=dtype).detach()
has_mask = True
if has_mask:
unmasked_target = get_noise_pred(
prompt_pair.positive_target, # negative prompt
prompt_pair.target_class, # positive prompt
1,
current_timestep,
denoised_latents
)
unmasked_target = unmasked_target.detach()
unmasked_target.requires_grad = False
else:
unmasked_target = None
# 4.20 GB RAM for 512x512
positive_latents = get_noise_pred(
positive, negative, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
prompt_pair.positive_target, # negative prompt
prompt_pair.negative_target, # positive prompt
1,
current_timestep,
denoised_latents
)
positive_latents = positive_latents.detach()
positive_latents.requires_grad = False
neutral_latents = get_noise_pred(
positive, neutral, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
prompt_pair.positive_target, # negative prompt
prompt_pair.empty_prompt, # positive prompt (normally neutral
1,
current_timestep,
denoised_latents
)
neutral_latents = neutral_latents.detach()
neutral_latents.requires_grad = False
unconditional_latents = get_noise_pred(
positive, positive, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
anchor_loss = None
if len(self.anchor_pairs) > 0:
# get a random anchor pair
anchor: EncodedAnchor = self.anchor_pairs[
torch.randint(0, len(self.anchor_pairs), (1,)).item()
]
with torch.no_grad():
anchor_target_noise = get_noise_pred(
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
with self.network:
# anchor whatever weight prompt pair is using
pos_nem_mult = 1.0 if prompt_pair.multiplier > 0 else -1.0
self.network.multiplier = anchor.multiplier * pos_nem_mult
anchor_pred_noise = get_noise_pred(
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
self.network.multiplier = prompt_pair.multiplier
with self.network:
self.network.multiplier = prompt_pair.multiplier
target_latents = get_noise_pred(
positive, target_class, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
# if self.logging_config.verbose:
# self.print("target_latents:", target_latents[0, 0, :5, :5])
positive_latents.requires_grad = False
neutral_latents.requires_grad = False
unconditional_latents.requires_grad = False
if len(self.anchor_pairs) > 0:
anchor_target_noise.requires_grad = False
anchor_loss = loss_function(
anchor_target_noise,
anchor_pred_noise,
prompt_pair.positive_target, # negative prompt
prompt_pair.positive_target, # positive prompt
1,
current_timestep,
denoised_latents
)
erase = prompt_pair.action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE
guidance_scale = 1.0
unconditional_latents = unconditional_latents.detach()
unconditional_latents.requires_grad = False
offset = guidance_scale * (positive_latents - unconditional_latents)
denoised_latents = denoised_latents.detach()
offset_neutral = neutral_latents
if erase:
offset_neutral -= offset
else:
# enhance
offset_neutral += offset
self.sd.set_device_state(self.train_slider_device_state)
self.sd.unet.train()
# start accumulating gradients
self.optimizer.zero_grad(set_to_none=True)
loss = loss_function(
target_latents,
offset_neutral,
) * weight
anchor_loss_float = None
if len(self.anchor_pairs) > 0:
with torch.no_grad():
# get a random anchor pair
anchor: EncodedAnchor = self.anchor_pairs[
torch.randint(0, len(self.anchor_pairs), (1,)).item()
]
anchor.to(self.device_torch, dtype=dtype)
loss_slide = loss.item()
# first we get the target prediction without network active
anchor_target_noise = get_noise_pred(
anchor.neg_prompt, anchor.prompt, 1, current_timestep, denoised_latents
# ).to("cpu", dtype=torch.float32)
).requires_grad_(False)
if anchor_loss is not None:
loss += anchor_loss
# to save vram, we will run these through separately while tracking grads
# otherwise it consumes a ton of vram and this isn't our speed bottleneck
anchor_chunks = split_anchors(anchor, self.prompt_chunk_size)
anchor_target_noise_chunks = torch.chunk(anchor_target_noise, self.prompt_chunk_size, dim=0)
assert len(anchor_chunks) == len(denoised_latent_chunks)
loss_float = loss.item()
# 4.32 GB RAM for 512x512
with self.network:
assert self.network.is_active
anchor_float_losses = []
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 + anchor_chunk.multiplier_list
loss = loss.to(self.device_torch)
anchor_pred_noise = get_noise_pred(
anchor_chunk.neg_prompt, anchor_chunk.prompt, 1, current_timestep, denoised_latent_chunk
)
# 9.42 GB RAM for 512x512 -> 4.20 GB RAM for 512x512 with new grad_checkpointing
anchor_loss = loss_function(
anchor_target_noise_chunk,
anchor_pred_noise,
)
anchor_float_losses.append(anchor_loss.item())
# compute anchor loss gradients
# we will accumulate them later
# this saves a ton of memory doing them separately
anchor_loss.backward()
del anchor_pred_noise
del anchor_target_noise_chunk
del anchor_loss
flush()
anchor_loss_float = sum(anchor_float_losses) / len(anchor_float_losses)
del anchor_chunks
del anchor_target_noise_chunks
del anchor_target_noise
# move anchor back to cpu
anchor.to("cpu")
with torch.no_grad():
if self.slider_config.low_ram:
prompt_pair_chunks = split_prompt_pairs(prompt_pair.detach(), self.prompt_chunk_size)
denoised_latent_chunks = denoised_latent_chunks # just to have it in one place
positive_latents_chunks = torch.chunk(positive_latents.detach(), self.prompt_chunk_size, dim=0)
neutral_latents_chunks = torch.chunk(neutral_latents.detach(), self.prompt_chunk_size, dim=0)
unconditional_latents_chunks = torch.chunk(
unconditional_latents.detach(),
self.prompt_chunk_size,
dim=0
)
mask_multiplier_chunks = torch.chunk(mask_multiplier, self.prompt_chunk_size, dim=0)
if unmasked_target is not None:
unmasked_target_chunks = torch.chunk(unmasked_target, self.prompt_chunk_size, dim=0)
else:
unmasked_target_chunks = [None for _ in range(self.prompt_chunk_size)]
else:
# run through in one instance
prompt_pair_chunks = [prompt_pair.detach()]
denoised_latent_chunks = [torch.cat(denoised_latent_chunks, dim=0).detach()]
positive_latents_chunks = [positive_latents.detach()]
neutral_latents_chunks = [neutral_latents.detach()]
unconditional_latents_chunks = [unconditional_latents.detach()]
mask_multiplier_chunks = [mask_multiplier]
unmasked_target_chunks = [unmasked_target]
# flush()
assert len(prompt_pair_chunks) == len(denoised_latent_chunks)
# 3.28 GB RAM for 512x512
with self.network:
assert self.network.is_active
loss_list = []
for prompt_pair_chunk, \
denoised_latent_chunk, \
positive_latents_chunk, \
neutral_latents_chunk, \
unconditional_latents_chunk, \
mask_multiplier_chunk, \
unmasked_target_chunk \
in zip(
prompt_pair_chunks,
denoised_latent_chunks,
positive_latents_chunks,
neutral_latents_chunks,
unconditional_latents_chunks,
mask_multiplier_chunks,
unmasked_target_chunks
):
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,
1,
current_timestep,
denoised_latent_chunk
)
guidance_scale = 1.0
offset = guidance_scale * (positive_latents_chunk - unconditional_latents_chunk)
# make offset multiplier based on actions
offset_multiplier_list = []
for action in prompt_pair_chunk.action_list:
if action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE:
offset_multiplier_list += [-1.0]
elif action == ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE:
offset_multiplier_list += [1.0]
offset_multiplier = torch.tensor(offset_multiplier_list).to(offset.device, dtype=offset.dtype)
# make offset multiplier match rank of offset
offset_multiplier = offset_multiplier.view(offset.shape[0], 1, 1, 1)
offset *= offset_multiplier
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")
# do inverted mask to preserve non masked
if has_mask and unmasked_target_chunk is not None:
loss = loss * mask_multiplier_chunk
# match the mask unmasked_target_chunk
mask_target_loss = torch.nn.functional.mse_loss(
target_latents.float(),
unmasked_target_chunk.float(),
reduction="none"
)
mask_target_loss = mask_target_loss * (1.0 - mask_multiplier_chunk)
loss += mask_target_loss
loss = loss.mean([1, 2, 3])
if self.train_config.learnable_snr_gos:
if from_batch:
# match batch size
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler,
self.train_config.min_snr_gamma)
else:
# match batch size
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
# add snr_gamma
loss = apply_learnable_snr_gos(loss, timesteps_index_list, self.snr_gos)
if self.train_config.min_snr_gamma is not None and self.train_config.min_snr_gamma > 0.000001:
if from_batch:
# match batch size
loss = apply_snr_weight(loss, timesteps, self.sd.noise_scheduler,
self.train_config.min_snr_gamma)
else:
# match batch size
timesteps_index_list = [current_timestep_index for _ in range(target_latents.shape[0])]
# add min_snr_gamma
loss = apply_snr_weight(loss, timesteps_index_list, noise_scheduler,
self.train_config.min_snr_gamma)
loss = loss.mean() * prompt_pair_chunk.weight
loss.backward()
loss_list.append(loss.item())
del target_latents
del offset_neutral
del loss
# flush()
loss.backward()
optimizer.step()
lr_scheduler.step()
loss_float = sum(loss_list) / len(loss_list)
if anchor_loss_float is not None:
loss_float += anchor_loss_float
del (
positive_latents,
neutral_latents,
unconditional_latents,
target_latents,
latents,
# latents
)
flush()
# move back to cpu
prompt_pair.to("cpu")
# flush()
# reset network
self.network.multiplier = 1.0
@@ -399,9 +686,9 @@ class TrainSliderProcess(BaseSDTrainProcess):
loss_dict = OrderedDict(
{'loss': loss_float},
)
if anchor_loss is not None:
loss_dict['sl_l'] = loss_slide
loss_dict['an_l'] = anchor_loss.item()
if anchor_loss_float is not None:
loss_dict['sl_l'] = loss_float
loss_dict['an_l'] = anchor_loss_float
return loss_dict
# end hook_train_loop

View File

@@ -0,0 +1,408 @@
# ref:
# - https://github.com/p1atdev/LECO/blob/main/train_lora.py
import time
from collections import OrderedDict
import os
from typing import Optional
from toolkit.config_modules import SliderConfig
from toolkit.paths import REPOS_ROOT
import sys
from toolkit.stable_diffusion_model import PromptEmbeds
sys.path.append(REPOS_ROOT)
sys.path.append(os.path.join(REPOS_ROOT, 'leco'))
from toolkit.train_tools import get_torch_dtype, apply_noise_offset
import gc
from toolkit import train_tools
import torch
from leco import train_util, model_util
from .BaseSDTrainProcess import BaseSDTrainProcess, StableDiffusion
class ACTION_TYPES_SLIDER:
ERASE_NEGATIVE = 0
ENHANCE_NEGATIVE = 1
def flush():
torch.cuda.empty_cache()
gc.collect()
class EncodedPromptPair:
def __init__(
self,
target_class,
positive,
negative,
neutral,
width=512,
height=512,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=1.0,
weight=1.0
):
self.target_class = target_class
self.positive = positive
self.negative = negative
self.neutral = neutral
self.width = width
self.height = height
self.action: int = action
self.multiplier = multiplier
self.weight = weight
class PromptEmbedsCache: # 使いまわしたいので
prompts: dict[str, PromptEmbeds] = {}
def __setitem__(self, __name: str, __value: PromptEmbeds) -> None:
self.prompts[__name] = __value
def __getitem__(self, __name: str) -> Optional[PromptEmbeds]:
if __name in self.prompts:
return self.prompts[__name]
else:
return None
class EncodedAnchor:
def __init__(
self,
prompt,
neg_prompt,
multiplier=1.0
):
self.prompt = prompt
self.neg_prompt = neg_prompt
self.multiplier = multiplier
class TrainSliderProcessOld(BaseSDTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
self.step_num = 0
self.start_step = 0
self.device = self.get_conf('device', self.job.device)
self.device_torch = torch.device(self.device)
self.slider_config = SliderConfig(**self.get_conf('slider', {}))
self.prompt_cache = PromptEmbedsCache()
self.prompt_pairs: list[EncodedPromptPair] = []
self.anchor_pairs: list[EncodedAnchor] = []
def before_model_load(self):
pass
def hook_before_train_loop(self):
cache = PromptEmbedsCache()
prompt_pairs: list[EncodedPromptPair] = []
# get encoded latents for our prompts
with torch.no_grad():
neutral = ""
for target in self.slider_config.targets:
# build the cache
for prompt in [
target.target_class,
target.positive,
target.negative,
neutral # empty neutral
]:
if cache[prompt] is None:
cache[prompt] = self.sd.encode_prompt(prompt)
for resolution in self.slider_config.resolutions:
width, height = resolution
only_erase = len(target.positive.strip()) == 0
only_enhance = len(target.negative.strip()) == 0
both = not only_erase and not only_enhance
if only_erase and only_enhance:
raise ValueError("target must have at least one of positive or negative or both")
# for slider we need to have an enhancer, an eraser, and then
# an inverse with negative weights to balance the network
# if we don't do this, we will get different contrast and focus.
# we only perform actions of enhancing and erasing on the negative
# todo work on way to do all of this in one shot
if both or only_erase:
prompt_pairs += [
# erase standard
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.positive],
negative=cache[target.negative],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=target.multiplier,
weight=target.weight
),
]
if both or only_enhance:
prompt_pairs += [
# enhance standard, swap pos neg
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.negative],
negative=cache[target.positive],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
multiplier=target.multiplier,
weight=target.weight
),
]
if both:
prompt_pairs += [
# erase inverted
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.negative],
negative=cache[target.positive],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ERASE_NEGATIVE,
multiplier=target.multiplier * -1.0,
weight=target.weight
),
]
prompt_pairs += [
# enhance inverted
EncodedPromptPair(
target_class=cache[target.target_class],
positive=cache[target.positive],
negative=cache[target.negative],
neutral=cache[neutral],
width=width,
height=height,
action=ACTION_TYPES_SLIDER.ENHANCE_NEGATIVE,
multiplier=target.multiplier * -1.0,
weight=target.weight
),
]
# setup anchors
anchor_pairs = []
for anchor in self.slider_config.anchors:
# build the cache
for prompt in [
anchor.prompt,
anchor.neg_prompt # empty neutral
]:
if cache[prompt] == None:
cache[prompt] = self.sd.encode_prompt(prompt)
anchor_pairs += [
EncodedAnchor(
prompt=cache[anchor.prompt],
neg_prompt=cache[anchor.neg_prompt],
multiplier=anchor.multiplier
)
]
# move to cpu to save vram
# We don't need text encoder anymore, but keep it on cpu for sampling
# if text encoder is list
if isinstance(self.sd.text_encoder, list):
for encoder in self.sd.text_encoder:
encoder.to("cpu")
else:
self.sd.text_encoder.to("cpu")
self.prompt_cache = cache
self.prompt_pairs = prompt_pairs
self.anchor_pairs = anchor_pairs
flush()
# end hook_before_train_loop
def hook_train_loop(self, batch):
dtype = get_torch_dtype(self.train_config.dtype)
# get a random pair
prompt_pair: EncodedPromptPair = self.prompt_pairs[
torch.randint(0, len(self.prompt_pairs), (1,)).item()
]
height = prompt_pair.height
width = prompt_pair.width
target_class = prompt_pair.target_class
neutral = prompt_pair.neutral
negative = prompt_pair.negative
positive = prompt_pair.positive
weight = prompt_pair.weight
multiplier = prompt_pair.multiplier
unet = self.sd.unet
noise_scheduler = self.sd.noise_scheduler
optimizer = self.optimizer
lr_scheduler = self.lr_scheduler
loss_function = torch.nn.MSELoss()
def get_noise_pred(p, n, gs, cts, dn):
return self.sd.predict_noise(
latents=dn,
text_embeddings=train_tools.concat_prompt_embeddings(
p, # unconditional
n, # positive
self.train_config.batch_size,
),
timestep=cts,
guidance_scale=gs,
)
# set network multiplier
self.network.multiplier = multiplier
with torch.no_grad():
self.sd.noise_scheduler.set_timesteps(
self.train_config.max_denoising_steps, device=self.device_torch
)
self.optimizer.zero_grad()
# ger a random number of steps
timesteps_to = torch.randint(
1, self.train_config.max_denoising_steps, (1,)
).item()
# get noise
noise = self.sd.get_latent_noise(
pixel_height=height,
pixel_width=width,
batch_size=self.train_config.batch_size,
noise_offset=self.train_config.noise_offset,
).to(self.device_torch, dtype=dtype)
# get latents
latents = noise * self.sd.noise_scheduler.init_noise_sigma
latents = latents.to(self.device_torch, dtype=dtype)
with self.network:
assert self.network.is_active
self.network.multiplier = multiplier
denoised_latents = self.sd.diffuse_some_steps(
latents, # pass simple noise latents
train_tools.concat_prompt_embeddings(
positive, # unconditional
target_class, # target
self.train_config.batch_size,
),
start_timesteps=0,
total_timesteps=timesteps_to,
guidance_scale=3,
)
noise_scheduler.set_timesteps(1000)
current_timestep = noise_scheduler.timesteps[
int(timesteps_to * 1000 / self.train_config.max_denoising_steps)
]
positive_latents = get_noise_pred(
positive, negative, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
neutral_latents = get_noise_pred(
positive, neutral, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
unconditional_latents = get_noise_pred(
positive, positive, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
anchor_loss = None
if len(self.anchor_pairs) > 0:
# get a random anchor pair
anchor: EncodedAnchor = self.anchor_pairs[
torch.randint(0, len(self.anchor_pairs), (1,)).item()
]
with torch.no_grad():
anchor_target_noise = get_noise_pred(
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
with self.network:
# anchor whatever weight prompt pair is using
pos_nem_mult = 1.0 if prompt_pair.multiplier > 0 else -1.0
self.network.multiplier = anchor.multiplier * pos_nem_mult
anchor_pred_noise = get_noise_pred(
anchor.prompt, anchor.neg_prompt, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
self.network.multiplier = prompt_pair.multiplier
with self.network:
self.network.multiplier = prompt_pair.multiplier
target_latents = get_noise_pred(
positive, target_class, 1, current_timestep, denoised_latents
).to("cpu", dtype=torch.float32)
# if self.logging_config.verbose:
# self.print("target_latents:", target_latents[0, 0, :5, :5])
positive_latents.requires_grad = False
neutral_latents.requires_grad = False
unconditional_latents.requires_grad = False
if len(self.anchor_pairs) > 0:
anchor_target_noise.requires_grad = False
anchor_loss = loss_function(
anchor_target_noise,
anchor_pred_noise,
)
erase = prompt_pair.action == ACTION_TYPES_SLIDER.ERASE_NEGATIVE
guidance_scale = 1.0
offset = guidance_scale * (positive_latents - unconditional_latents)
offset_neutral = neutral_latents
if erase:
offset_neutral -= offset
else:
# enhance
offset_neutral += offset
loss = loss_function(
target_latents,
offset_neutral,
) * weight
loss_slide = loss.item()
if anchor_loss is not None:
loss += anchor_loss
loss_float = loss.item()
loss = loss.to(self.device_torch)
loss.backward()
optimizer.step()
lr_scheduler.step()
del (
positive_latents,
neutral_latents,
unconditional_latents,
target_latents,
latents,
)
flush()
# reset network
self.network.multiplier = 1.0
loss_dict = OrderedDict(
{'loss': loss_float},
)
if anchor_loss is not None:
loss_dict['sl_l'] = loss_slide
loss_dict['an_l'] = anchor_loss.item()
return loss_dict
# end hook_train_loop

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
@@ -24,6 +26,9 @@ from diffusers import AutoencoderKL
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(
[
@@ -37,145 +42,6 @@ def unnormalize(tensor):
return (tensor / 2 + 0.5).clamp(0, 1)
class Critic:
process: 'TrainVAEProcess'
def __init__(
self,
learning_rate=1e-5,
device='cpu',
optimizer='adam',
num_critic_per_gen=1,
dtype='float32',
lambda_gp=10,
start_step=0,
warmup_steps=1000,
process=None,
optimizer_params=None,
):
self.learning_rate = learning_rate
self.device = device
self.optimizer_type = optimizer
self.num_critic_per_gen = num_critic_per_gen
self.dtype = dtype
self.torch_dtype = get_torch_dtype(self.dtype)
self.process = process
self.model = None
self.optimizer = None
self.scheduler = None
self.warmup_steps = warmup_steps
self.start_step = start_step
self.lambda_gp = lambda_gp
if optimizer_params is None:
optimizer_params = {}
self.optimizer_params = optimizer_params
self.print = self.process.print
print(f" Critic config: {self.__dict__}")
def setup(self):
from .models.vgg19_critic import Vgg19Critic
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
self.load_weights()
self.model.train()
self.model.requires_grad_(True)
params = self.model.parameters()
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
optimizer_params=self.optimizer_params)
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
self.optimizer,
total_iters=self.process.max_steps * self.num_critic_per_gen,
factor=1,
verbose=False
)
def load_weights(self):
path_to_load = None
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
if files and len(files) > 0:
latest_file = max(files, key=os.path.getmtime)
print(f" - Latest checkpoint is: {latest_file}")
path_to_load = latest_file
else:
self.print(f" - No checkpoint found, starting from scratch")
if path_to_load:
self.model.load_state_dict(load_file(path_to_load))
def save(self, step=None):
self.process.update_training_metadata()
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
step_num = ''
if step is not None:
# zeropad 9 digits
step_num = f"_{str(step).zfill(9)}"
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
save_file(self.model.state_dict(), save_path, save_meta)
self.print(f"Saved critic to {save_path}")
def get_critic_loss(self, vgg_output):
if self.start_step > self.process.step_num:
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
warmup_scaler = 1.0
# we need a warmup when we come on of 1000 steps
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
if self.process.step_num < self.start_step + self.warmup_steps:
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
# set model to not train for generator loss
self.model.eval()
self.model.requires_grad_(False)
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
# run model
stacked_output = self.model(vgg_pred)
return (-torch.mean(stacked_output)) * warmup_scaler
def step(self, vgg_output):
# train critic here
self.model.train()
self.model.requires_grad_(True)
critic_losses = []
for i in range(self.num_critic_per_gen):
inputs = vgg_output.detach()
inputs = inputs.to(self.device, dtype=self.torch_dtype)
self.optimizer.zero_grad()
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
stacked_output = self.model(inputs)
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
# Compute gradient penalty
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
# Compute WGAN-GP critic loss
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
critic_loss.backward()
self.optimizer.zero_grad()
self.optimizer.step()
self.scheduler.step()
critic_losses.append(critic_loss.item())
# avg loss
loss = np.mean(critic_losses)
return loss
def get_lr(self):
if self.optimizer_type.startswith('dadaptation'):
learning_rate = (
self.optimizer.param_groups[0]["d"] *
self.optimizer.param_groups[0]["lr"]
)
else:
learning_rate = self.optimizer.param_groups[0]['lr']
return learning_rate
class TrainVAEProcess(BaseTrainProcess):
def __init__(self, process_id: int, job, config: OrderedDict):
super().__init__(process_id, job, config)
@@ -200,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', {})
@@ -209,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
@@ -275,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(
@@ -349,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
@@ -364,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):
@@ -418,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
@@ -432,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}")
@@ -444,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()
@@ -512,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)
@@ -535,6 +422,7 @@ class TrainVAEProcess(BaseTrainProcess):
self.sample()
blank_losses = OrderedDict({
"total": [],
"lpips": [],
"style": [],
"content": [],
"mse": [],
@@ -553,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)
@@ -579,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()
@@ -598,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:
@@ -615,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"]
@@ -633,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())
@@ -643,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

@@ -6,5 +6,10 @@ from .BaseTrainProcess import BaseTrainProcess
from .TrainVAEProcess import TrainVAEProcess
from .BaseMergeProcess import BaseMergeProcess
from .TrainSliderProcess import TrainSliderProcess
from .TrainLoRAHack import TrainLoRAHack
from .TrainSDRescaleProcess import TrainSDRescaleProcess
from .TrainSliderProcessOld import TrainSliderProcessOld
from .TrainSDRescaleProcess import TrainSDRescaleProcess
from .ModRescaleLoraProcess import ModRescaleLoraProcess
from .GenerateProcess import GenerateProcess
from .BaseExtensionProcess import BaseExtensionProcess
from .TrainESRGANProcess import TrainESRGANProcess
from .BaseSDTrainProcess import BaseSDTrainProcess

View File

@@ -1,5 +1,17 @@
import glob
import os
import numpy as np
import torch
import torch.nn as nn
from safetensors.torch import load_file, save_file
from toolkit.losses import get_gradient_penalty
from toolkit.metadata import get_meta_for_safetensors
from toolkit.optimizer import get_optimizer
from toolkit.train_tools import get_torch_dtype
from typing import TYPE_CHECKING, Union
class MeanReduce(nn.Module):
@@ -36,3 +48,147 @@ class Vgg19Critic(nn.Module):
def forward(self, inputs):
return self.main(inputs)
if TYPE_CHECKING:
from jobs.process.TrainVAEProcess import TrainVAEProcess
from jobs.process.TrainESRGANProcess import TrainESRGANProcess
class Critic:
process: Union['TrainVAEProcess', 'TrainESRGANProcess']
def __init__(
self,
learning_rate=1e-5,
device='cpu',
optimizer='adam',
num_critic_per_gen=1,
dtype='float32',
lambda_gp=10,
start_step=0,
warmup_steps=1000,
process=None,
optimizer_params=None,
):
self.learning_rate = learning_rate
self.device = device
self.optimizer_type = optimizer
self.num_critic_per_gen = num_critic_per_gen
self.dtype = dtype
self.torch_dtype = get_torch_dtype(self.dtype)
self.process = process
self.model = None
self.optimizer = None
self.scheduler = None
self.warmup_steps = warmup_steps
self.start_step = start_step
self.lambda_gp = lambda_gp
if optimizer_params is None:
optimizer_params = {}
self.optimizer_params = optimizer_params
self.print = self.process.print
print(f" Critic config: {self.__dict__}")
def setup(self):
self.model = Vgg19Critic().to(self.device, dtype=self.torch_dtype)
self.load_weights()
self.model.train()
self.model.requires_grad_(True)
params = self.model.parameters()
self.optimizer = get_optimizer(params, self.optimizer_type, self.learning_rate,
optimizer_params=self.optimizer_params)
self.scheduler = torch.optim.lr_scheduler.ConstantLR(
self.optimizer,
total_iters=self.process.max_steps * self.num_critic_per_gen,
factor=1,
verbose=False
)
def load_weights(self):
path_to_load = None
self.print(f"Critic: Looking for latest checkpoint in {self.process.save_root}")
files = glob.glob(os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}*.safetensors"))
if files and len(files) > 0:
latest_file = max(files, key=os.path.getmtime)
print(f" - Latest checkpoint is: {latest_file}")
path_to_load = latest_file
else:
self.print(f" - No checkpoint found, starting from scratch")
if path_to_load:
self.model.load_state_dict(load_file(path_to_load))
def save(self, step=None):
self.process.update_training_metadata()
save_meta = get_meta_for_safetensors(self.process.meta, self.process.job.name)
step_num = ''
if step is not None:
# zeropad 9 digits
step_num = f"_{str(step).zfill(9)}"
save_path = os.path.join(self.process.save_root, f"CRITIC_{self.process.job.name}{step_num}.safetensors")
save_file(self.model.state_dict(), save_path, save_meta)
self.print(f"Saved critic to {save_path}")
def get_critic_loss(self, vgg_output):
if self.start_step > self.process.step_num:
return torch.tensor(0.0, dtype=self.torch_dtype, device=self.device)
warmup_scaler = 1.0
# we need a warmup when we come on of 1000 steps
# we want to scale the loss by 0.0 at self.start_step steps and 1.0 at self.start_step + warmup_steps
if self.process.step_num < self.start_step + self.warmup_steps:
warmup_scaler = (self.process.step_num - self.start_step) / self.warmup_steps
# set model to not train for generator loss
self.model.eval()
self.model.requires_grad_(False)
vgg_pred, vgg_target = torch.chunk(vgg_output, 2, dim=0)
# run model
stacked_output = self.model(vgg_pred)
return (-torch.mean(stacked_output)) * warmup_scaler
def step(self, vgg_output):
# train critic here
self.model.train()
self.model.requires_grad_(True)
self.optimizer.zero_grad()
critic_losses = []
inputs = vgg_output.detach()
inputs = inputs.to(self.device, dtype=self.torch_dtype)
self.optimizer.zero_grad()
vgg_pred, vgg_target = torch.chunk(inputs, 2, dim=0)
stacked_output = self.model(inputs).float()
out_pred, out_target = torch.chunk(stacked_output, 2, dim=0)
# Compute gradient penalty
gradient_penalty = get_gradient_penalty(self.model, vgg_target, vgg_pred, self.device)
# Compute WGAN-GP critic loss
critic_loss = -(torch.mean(out_target) - torch.mean(out_pred)) + self.lambda_gp * gradient_penalty
critic_loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)
self.optimizer.step()
self.scheduler.step()
critic_losses.append(critic_loss.item())
# avg loss
loss = np.mean(critic_losses)
return loss
def get_lr(self):
if self.optimizer_type.startswith('dadaptation'):
learning_rate = (
self.optimizer.param_groups[0]["d"] *
self.optimizer.param_groups[0]["lr"]
)
else:
learning_rate = self.optimizer.param_groups[0]['lr']
return learning_rate

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

@@ -0,0 +1,339 @@
{
"nbformat": 4,
"nbformat_minor": 0,
"metadata": {
"colab": {
"provenance": [],
"machine_shape": "hm",
"gpuType": "V100"
},
"kernelspec": {
"name": "python3",
"display_name": "Python 3"
},
"language_info": {
"name": "python"
},
"accelerator": "GPU"
},
"cells": [
{
"cell_type": "markdown",
"source": [
"# AI Toolkit by Ostris\n",
"## Slider Training\n",
"\n",
"This is a quick colab demo for training sliders like can be found in my CivitAI profile https://civitai.com/user/Ostris/models . I will work on making it more user friendly, but for now, it will get you started."
],
"metadata": {
"collapsed": false
}
},
{
"cell_type": "code",
"source": [
"!git clone https://github.com/ostris/ai-toolkit"
],
"metadata": {
"id": "BvAG0GKAh59G"
},
"execution_count": null,
"outputs": []
},
{
"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": "code",
"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"
],
"metadata": {
"collapsed": false
},
"outputs": []
},
{
"cell_type": "markdown",
"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."
],
"metadata": {
"id": "N8UUFzVRigbC"
}
},
{
"cell_type": "code",
"source": [
"from collections import OrderedDict\n",
"\n",
"job_to_run = OrderedDict({\n",
" # This is the config I use on my sliders, It is solid and tested\n",
" 'job': 'train',\n",
" 'config': {\n",
" # the name will be used to create a folder in the output folder\n",
" # it will also replace any [name] token in the rest of this config\n",
" 'name': 'detail_slider_v1',\n",
" # folder will be created with name above in folder below\n",
" # it can be relative to the project root or absolute\n",
" 'training_folder': \"output/LoRA\",\n",
" 'device': 'cuda', # cpu, cuda:0, etc\n",
" # for tensorboard logging, we will make a subfolder for this job\n",
" 'log_dir': \"output/.tensorboard\",\n",
" # you can stack processes for other jobs, It is not tested with sliders though\n",
" # just use one for now\n",
" 'process': [\n",
" {\n",
" 'type': 'slider', # tells runner to run the slider process\n",
" # network is the LoRA network for a slider, I recommend to leave this be\n",
" 'network': {\n",
" 'type': \"lora\",\n",
" # rank / dim of the network. Bigger is not always better. Especially for sliders. 8 is good\n",
" 'linear': 8, # \"rank\" or \"dim\"\n",
" 'linear_alpha': 4, # Do about half of rank \"alpha\"\n",
" # 'conv': 4, # for convolutional layers \"locon\"\n",
" # 'conv_alpha': 4, # Do about half of conv \"alpha\"\n",
" },\n",
" # training config\n",
" 'train': {\n",
" # this is also used in sampling. Stick with ddpm unless you know what you are doing\n",
" 'noise_scheduler': \"ddpm\", # or \"ddpm\", \"lms\", \"euler_a\"\n",
" # how many steps to train. More is not always better. I rarely go over 1000\n",
" 'steps': 100,\n",
" # I have had good results with 4e-4 to 1e-4 at 500 steps\n",
" 'lr': 2e-4,\n",
" # enables gradient checkpoint, saves vram, leave it on\n",
" 'gradient_checkpointing': True,\n",
" # train the unet. I recommend leaving this true\n",
" 'train_unet': True,\n",
" # train the text encoder. I don't recommend this unless you have a special use case\n",
" # for sliders we are adjusting representation of the concept (unet),\n",
" # not the description of it (text encoder)\n",
" 'train_text_encoder': False,\n",
"\n",
" # just leave unless you know what you are doing\n",
" # also supports \"dadaptation\" but set lr to 1 if you use that,\n",
" # but it learns too fast and I don't recommend it\n",
" 'optimizer': \"adamw\",\n",
" # only constant for now\n",
" 'lr_scheduler': \"constant\",\n",
" # we randomly denoise random num of steps form 1 to this number\n",
" # while training. Just leave it\n",
" 'max_denoising_steps': 40,\n",
" # works great at 1. I do 1 even with my 4090.\n",
" # higher may not work right with newer single batch stacking code anyway\n",
" 'batch_size': 1,\n",
" # bf16 works best if your GPU supports it (modern)\n",
" 'dtype': 'bf16', # fp32, bf16, fp16\n",
" # I don't recommend using unless you are trying to make a darker lora. Then do 0.1 MAX\n",
" # although, the way we train sliders is comparative, so it probably won't work anyway\n",
" 'noise_offset': 0.0,\n",
" },\n",
"\n",
" # the model to train the LoRA network on\n",
" 'model': {\n",
" # name_or_path can be a hugging face name, local path or url to model\n",
" # on civit ai with or without modelVersionId. They will be cached in /model folder\n",
" # epicRealisim v5\n",
" 'name_or_path': \"https://civitai.com/models/25694?modelVersionId=134065\",\n",
" 'is_v2': False, # for v2 models\n",
" 'is_v_pred': False, # for v-prediction models (most v2 models)\n",
" # has some issues with the dual text encoder and the way we train sliders\n",
" # it works bit weights need to probably be higher to see it.\n",
" 'is_xl': False, # for SDXL models\n",
" },\n",
"\n",
" # saving config\n",
" 'save': {\n",
" 'dtype': 'float16', # precision to save. I recommend float16\n",
" 'save_every': 50, # save every this many steps\n",
" # this will remove step counts more than this number\n",
" # allows you to save more often in case of a crash without filling up your drive\n",
" 'max_step_saves_to_keep': 2,\n",
" },\n",
"\n",
" # sampling config\n",
" 'sample': {\n",
" # must match train.noise_scheduler, this is not used here\n",
" # but may be in future and in other processes\n",
" 'sampler': \"ddpm\",\n",
" # sample every this many steps\n",
" 'sample_every': 20,\n",
" # image size\n",
" 'width': 512,\n",
" 'height': 512,\n",
" # prompts to use for sampling. Do as many as you want, but it slows down training\n",
" # pick ones that will best represent the concept you are trying to adjust\n",
" # allows some flags after the prompt\n",
" # --m [number] # network multiplier. LoRA weight. -3 for the negative slide, 3 for the positive\n",
" # slide are good tests. will inherit sample.network_multiplier if not set\n",
" # --n [string] # negative prompt, will inherit sample.neg if not set\n",
" # Only 75 tokens allowed currently\n",
" # I like to do a wide positive and negative spread so I can see a good range and stop\n",
" # early if the network is braking down\n",
" 'prompts': [\n",
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -5\",\n",
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m -3\",\n",
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 3\",\n",
" \"a woman in a coffee shop, black hat, blonde hair, blue jacket --m 5\",\n",
" \"a golden retriever sitting on a leather couch, --m -5\",\n",
" \"a golden retriever sitting on a leather couch --m -3\",\n",
" \"a golden retriever sitting on a leather couch --m 3\",\n",
" \"a golden retriever sitting on a leather couch --m 5\",\n",
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -5\",\n",
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m -3\",\n",
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 3\",\n",
" \"a man with a beard and red flannel shirt, wearing vr goggles, walking into traffic --m 5\",\n",
" ],\n",
" # negative prompt used on all prompts above as default if they don't have one\n",
" 'neg': \"cartoon, fake, drawing, illustration, cgi, animated, anime, monochrome\",\n",
" # seed for sampling. 42 is the answer for everything\n",
" 'seed': 42,\n",
" # walks the seed so s1 is 42, s2 is 43, s3 is 44, etc\n",
" # will start over on next sample_every so s1 is always seed\n",
" # works well if you use same prompt but want different results\n",
" 'walk_seed': False,\n",
" # cfg scale (4 to 10 is good)\n",
" 'guidance_scale': 7,\n",
" # sampler steps (20 to 30 is good)\n",
" 'sample_steps': 20,\n",
" # default network multiplier for all prompts\n",
" # since we are training a slider, I recommend overriding this with --m [number]\n",
" # in the prompts above to get both sides of the slider\n",
" 'network_multiplier': 1.0,\n",
" },\n",
"\n",
" # logging information\n",
" 'logging': {\n",
" 'log_every': 10, # log every this many steps\n",
" 'use_wandb': False, # not supported yet\n",
" 'verbose': False, # probably done need unless you are debugging\n",
" },\n",
"\n",
" # slider training config, best for last\n",
" 'slider': {\n",
" # resolutions to train on. [ width, height ]. This is less important for sliders\n",
" # as we are not teaching the model anything it doesn't already know\n",
" # but must be a size it understands [ 512, 512 ] for sd_v1.5 and [ 768, 768 ] for sd_v2.1\n",
" # and [ 1024, 1024 ] for sd_xl\n",
" # you can do as many as you want here\n",
" 'resolutions': [\n",
" [512, 512],\n",
" # [ 512, 768 ]\n",
" # [ 768, 768 ]\n",
" ],\n",
" # slider training uses 4 combined steps for a single round. This will do it in one gradient\n",
" # step. It is highly optimized and shouldn't take anymore vram than doing without it,\n",
" # since we break down batches for gradient accumulation now. so just leave it on.\n",
" 'batch_full_slide': True,\n",
" # These are the concepts to train on. You can do as many as you want here,\n",
" # but they can conflict outweigh each other. Other than experimenting, I recommend\n",
" # just doing one for good results\n",
" 'targets': [\n",
" # target_class is the base concept we are adjusting the representation of\n",
" # for example, if we are adjusting the representation of a person, we would use \"person\"\n",
" # if we are adjusting the representation of a cat, we would use \"cat\" It is not\n",
" # a keyword necessarily but what the model understands the concept to represent.\n",
" # \"person\" will affect men, women, children, etc but will not affect cats, dogs, etc\n",
" # it is the models base general understanding of the concept and everything it represents\n",
" # you can leave it blank to affect everything. In this example, we are adjusting\n",
" # detail, so we will leave it blank to affect everything\n",
" {\n",
" 'target_class': \"\",\n",
" # positive is the prompt for the positive side of the slider.\n",
" # It is the concept that will be excited and amplified in the model when we slide the slider\n",
" # to the positive side and forgotten / inverted when we slide\n",
" # the slider to the negative side. It is generally best to include the target_class in\n",
" # the prompt. You want it to be the extreme of what you want to train on. For example,\n",
" # if you want to train on fat people, you would use \"an extremely fat, morbidly obese person\"\n",
" # as the prompt. Not just \"fat person\"\n",
" # max 75 tokens for now\n",
" 'positive': \"high detail, 8k, intricate, detailed, high resolution, high res, high quality\",\n",
" # negative is the prompt for the negative side of the slider and works the same as positive\n",
" # it does not necessarily work the same as a negative prompt when generating images\n",
" # these need to be polar opposites.\n",
" # max 76 tokens for now\n",
" 'negative': \"blurry, boring, fuzzy, low detail, low resolution, low res, low quality\",\n",
" # the loss for this target is multiplied by this number.\n",
" # if you are doing more than one target it may be good to set less important ones\n",
" # to a lower number like 0.1 so they don't outweigh the primary target\n",
" 'weight': 1.0,\n",
" },\n",
" ],\n",
" },\n",
" },\n",
" ]\n",
" },\n",
"\n",
" # You can put any information you want here, and it will be saved in the model.\n",
" # The below is an example, but you can put your grocery list in it if you want.\n",
" # It is saved in the model so be aware of that. The software will include this\n",
" # plus some other information for you automatically\n",
" 'meta': {\n",
" # [name] gets replaced with the name above\n",
" 'name': \"[name]\",\n",
" 'version': '1.0',\n",
" # 'creator': {\n",
" # 'name': 'your name',\n",
" # 'email': 'your@gmail.com',\n",
" # 'website': 'https://your.website'\n",
" # }\n",
" }\n",
"})\n"
],
"metadata": {
"id": "_t28QURYjRQO"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "markdown",
"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. Ill update soon."
],
"metadata": {
"id": "h6F1FlM2Wb3l"
}
},
{
"cell_type": "code",
"source": [
"run_job(job_to_run)\n"
],
"metadata": {
"id": "HkajwI8gteOh"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "markdown",
"source": [
"## Done\n",
"\n",
"Check your ourput dir and get your slider\n"
],
"metadata": {
"id": "Hblgb5uwW5SD"
}
}
]
}

View File

@@ -1,11 +1,10 @@
torch
torchvision
safetensors
diffusers
git+https://github.com/huggingface/diffusers.git
transformers
lycoris_lora
lycoris-lora==1.8.3
flatten_json
accelerator
pyyaml
oyaml
tensorboard
@@ -14,4 +13,20 @@ invisible-watermark
einops
accelerate
toml
albumentations
albumentations
pydantic
omegaconf
k-diffusion
open_clip_torch
timm
prodigyopt
controlnet_aux==0.0.7
python-dotenv
bitsandbytes
hf_transfer
lpips
pytorch_fid
optimum-quanto
sentencepiece
huggingface_hub
peft

27
run.py
View File

@@ -1,6 +1,23 @@
import os
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
import sys
from typing import Union, OrderedDict
from dotenv import load_dotenv
# Load the .env file if it exists
load_dotenv()
sys.path.insert(0, os.getcwd())
# 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'
# 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
@@ -36,6 +53,14 @@ def main():
action='store_true',
help='Continue running additional jobs even if a job fails'
)
# flag to continue if failed job
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()
config_file_list = args.config_file_list
@@ -49,7 +74,7 @@ def main():
for config_file in config_file_list:
try:
job = get_job(config_file)
job = get_job(config_file, args.name)
job.run()
job.cleanup()
jobs_completed += 1

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)

128
scripts/convert_cog.py Normal file
View File

@@ -0,0 +1,128 @@
import json
from collections import OrderedDict
import os
import torch
from safetensors import safe_open
from safetensors.torch import save_file
device = torch.device('cpu')
# [diffusers] -> kohya
embedding_mapping = {
'text_encoders_0': 'clip_l',
'text_encoders_1': 'clip_g'
}
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
KEYMAP_ROOT = os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps')
sdxl_keymap_path = os.path.join(KEYMAP_ROOT, 'stable_diffusion_locon_sdxl.json')
# load keymap
with open(sdxl_keymap_path, 'r') as f:
ldm_diffusers_keymap = json.load(f)['ldm_diffusers_keymap']
# invert the item / key pairs
diffusers_ldm_keymap = {v: k for k, v in ldm_diffusers_keymap.items()}
def get_ldm_key(diffuser_key):
diffuser_key = f"lora_unet_{diffuser_key.replace('.', '_')}"
diffuser_key = diffuser_key.replace('_lora_down_weight', '.lora_down.weight')
diffuser_key = diffuser_key.replace('_lora_up_weight', '.lora_up.weight')
diffuser_key = diffuser_key.replace('_alpha', '.alpha')
diffuser_key = diffuser_key.replace('_processor_to_', '_to_')
diffuser_key = diffuser_key.replace('_to_out.', '_to_out_0.')
if diffuser_key in diffusers_ldm_keymap:
return diffusers_ldm_keymap[diffuser_key]
else:
raise KeyError(f"Key {diffuser_key} not found in keymap")
def convert_cog(lora_path, embedding_path):
embedding_state_dict = OrderedDict()
lora_state_dict = OrderedDict()
# # normal dict
# normal_dict = OrderedDict()
# example_path = "/mnt/Models/stable-diffusion/models/LoRA/sdxl/LogoRedmond_LogoRedAF.safetensors"
# with safe_open(example_path, framework="pt", device='cpu') as f:
# keys = list(f.keys())
# for key in keys:
# normal_dict[key] = f.get_tensor(key)
with safe_open(embedding_path, framework="pt", device='cpu') as f:
keys = list(f.keys())
for key in keys:
new_key = embedding_mapping[key]
embedding_state_dict[new_key] = f.get_tensor(key)
with safe_open(lora_path, framework="pt", device='cpu') as f:
keys = list(f.keys())
lora_rank = None
# get the lora dim first. Check first 3 linear layers just to be safe
for key in keys:
new_key = get_ldm_key(key)
tensor = f.get_tensor(key)
num_checked = 0
if len(tensor.shape) == 2:
this_dim = min(tensor.shape)
if lora_rank is None:
lora_rank = this_dim
elif lora_rank != this_dim:
raise ValueError(f"lora rank is not consistent, got {tensor.shape}")
else:
num_checked += 1
if num_checked >= 3:
break
for key in keys:
new_key = get_ldm_key(key)
tensor = f.get_tensor(key)
if new_key.endswith('.lora_down.weight'):
alpha_key = new_key.replace('.lora_down.weight', '.alpha')
# diffusers does not have alpha, they usa an alpha multiplier of 1 which is a tensor weight of the dims
# assume first smallest dim is the lora rank if shape is 2
lora_state_dict[alpha_key] = torch.ones(1).to(tensor.device, tensor.dtype) * lora_rank
lora_state_dict[new_key] = tensor
return lora_state_dict, embedding_state_dict
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument(
'lora_path',
type=str,
help='Path to lora file'
)
parser.add_argument(
'embedding_path',
type=str,
help='Path to embedding file'
)
parser.add_argument(
'--lora_output',
type=str,
default="lora_output",
)
parser.add_argument(
'--embedding_output',
type=str,
default="embedding_output",
)
args = parser.parse_args()
lora_state_dict, embedding_state_dict = convert_cog(args.lora_path, args.embedding_path)
# save them
save_file(lora_state_dict, args.lora_output)
save_file(embedding_state_dict, args.embedding_output)
print(f"Saved lora to {args.lora_output}")
print(f"Saved embedding to {args.embedding_output}")

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

@@ -0,0 +1,61 @@
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
from toolkit.config_modules import ModelConfig
from toolkit.stable_diffusion_model import StableDiffusion
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'
)
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('--sd2', action='store_true', help='is sd 2 model')
args = parser.parse_args()
device = torch.device('cpu')
dtype = torch.float32
print(f"Loading model from {args.input_path}")
diffusers_model_config = ModelConfig(
name_or_path=args.input_path,
is_xl=args.sdxl,
is_v2=args.sd2,
is_ssd=args.ssd,
dtype=dtype,
)
diffusers_sd = StableDiffusion(
model_config=diffusers_model_config,
device=device,
dtype=dtype,
)
diffusers_sd.load_model()
print(f"Loaded model from {args.input_path}")
diffusers_sd.pipeline.fuse_lora()
meta = OrderedDict()
diffusers_sd.save(args.output_path, meta=meta)
print(f"Saved to {args.output_path}")

View File

@@ -0,0 +1,67 @@
import argparse
from collections import OrderedDict
import torch
from toolkit.config_modules import ModelConfig
from toolkit.stable_diffusion_model import StableDiffusion
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'
)
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('--sd2', action='store_true', help='is sd 2 model')
args = parser.parse_args()
device = torch.device('cpu')
dtype = torch.float32
print(f"Loading model from {args.input_path}")
if args.sdxl:
adapter_id = "latent-consistency/lcm-lora-sdxl"
if args.refiner:
adapter_id = "latent-consistency/lcm-lora-sdxl"
elif args.ssd:
adapter_id = "latent-consistency/lcm-lora-ssd-1b"
else:
adapter_id = "latent-consistency/lcm-lora-sdv1-5"
diffusers_model_config = ModelConfig(
name_or_path=args.input_path,
is_xl=args.sdxl,
is_v2=args.sd2,
is_ssd=args.ssd,
dtype=dtype,
)
diffusers_sd = StableDiffusion(
model_config=diffusers_model_config,
device=device,
dtype=dtype,
)
diffusers_sd.load_model()
print(f"Loaded model from {args.input_path}")
diffusers_sd.pipeline.load_lora_weights(adapter_id)
diffusers_sd.pipeline.fuse_lora()
meta = OrderedDict()
diffusers_sd.save(args.output_path, meta=meta)
print(f"Saved to {args.output_path}")

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

View File

@@ -1,547 +0,0 @@
import gc
import time
import argparse
import itertools
import math
import os
from multiprocessing import Value
from tqdm import tqdm
import torch
from accelerate.utils import set_seed
import diffusers
from diffusers import DDPMScheduler
import library.train_util as train_util
import library.config_util as config_util
from library.config_util import (
ConfigSanitizer,
BlueprintGenerator,
)
import custom_tools.train_tools as train_tools
import library.custom_train_functions as custom_train_functions
from library.custom_train_functions import (
apply_snr_weight,
get_weighted_text_embeddings,
prepare_scheduler_for_custom_training,
pyramid_noise_like,
apply_noise_offset,
scale_v_prediction_loss_like_noise_prediction,
)
# perlin_noise,
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
SD_SCRIPTS_ROOT = os.path.join(PROJECT_ROOT, "repositories", "sd-scripts")
def train(args):
train_util.verify_training_args(args)
train_util.prepare_dataset_args(args, False)
cache_latents = args.cache_latents
if args.seed is not None:
set_seed(args.seed) # 乱数系列を初期化する
tokenizer = train_util.load_tokenizer(args)
# データセットを準備する
if args.dataset_class is None:
blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, False, True))
if args.dataset_config is not None:
print(f"Load dataset config from {args.dataset_config}")
user_config = config_util.load_user_config(args.dataset_config)
ignored = ["train_data_dir", "reg_data_dir"]
if any(getattr(args, attr) is not None for attr in ignored):
print(
"ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format(
", ".join(ignored)
)
)
else:
user_config = {
"datasets": [
{"subsets": config_util.generate_dreambooth_subsets_config_by_subdirs(args.train_data_dir, args.reg_data_dir)}
]
}
blueprint = blueprint_generator.generate(user_config, args, tokenizer=tokenizer)
train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group)
else:
train_dataset_group = train_util.load_arbitrary_dataset(args, tokenizer)
current_epoch = Value("i", 0)
current_step = Value("i", 0)
ds_for_collater = train_dataset_group if args.max_data_loader_n_workers == 0 else None
collater = train_util.collater_class(current_epoch, current_step, ds_for_collater)
if args.no_token_padding:
train_dataset_group.disable_token_padding()
if args.debug_dataset:
train_util.debug_dataset(train_dataset_group)
return
if cache_latents:
assert (
train_dataset_group.is_latent_cacheable()
), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません"
# replace captions with names
if args.name_replace is not None:
print(f"Replacing captions [name] with '{args.name_replace}'")
train_dataset_group = train_tools.replace_filewords_in_dataset_group(
train_dataset_group, args
)
# acceleratorを準備する
print("prepare accelerator")
if args.gradient_accumulation_steps > 1:
print(
f"gradient_accumulation_steps is {args.gradient_accumulation_steps}. accelerate does not support gradient_accumulation_steps when training multiple models (U-Net and Text Encoder), so something might be wrong"
)
print(
f"gradient_accumulation_stepsが{args.gradient_accumulation_steps}に設定されています。accelerateは複数モデル(U-NetおよびText Encoder)の学習時にgradient_accumulation_stepsをサポートしていないため結果は未知数です"
)
accelerator, unwrap_model = train_util.prepare_accelerator(args)
# mixed precisionに対応した型を用意しておき適宜castする
weight_dtype, save_dtype = train_util.prepare_dtype(args)
# モデルを読み込む
text_encoder, vae, unet, load_stable_diffusion_format = train_util.load_target_model(args, weight_dtype, accelerator)
# verify load/save model formats
if load_stable_diffusion_format:
src_stable_diffusion_ckpt = args.pretrained_model_name_or_path
src_diffusers_model_path = None
else:
src_stable_diffusion_ckpt = None
src_diffusers_model_path = args.pretrained_model_name_or_path
if args.save_model_as is None:
save_stable_diffusion_format = load_stable_diffusion_format
use_safetensors = args.use_safetensors
else:
save_stable_diffusion_format = args.save_model_as.lower() == "ckpt" or args.save_model_as.lower() == "safetensors"
use_safetensors = args.use_safetensors or ("safetensors" in args.save_model_as.lower())
# モデルに xformers とか memory efficient attention を組み込む
train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers)
# 学習を準備する
if cache_latents:
vae.to(accelerator.device, dtype=weight_dtype)
vae.requires_grad_(False)
vae.eval()
with torch.no_grad():
train_dataset_group.cache_latents(vae, args.vae_batch_size, args.cache_latents_to_disk, accelerator.is_main_process)
vae.to("cpu")
if torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
accelerator.wait_for_everyone()
# 学習を準備する:モデルを適切な状態にする
train_text_encoder = args.stop_text_encoder_training is None or args.stop_text_encoder_training >= 0
unet.requires_grad_(True) # 念のため追加
text_encoder.requires_grad_(train_text_encoder)
if not train_text_encoder:
print("Text Encoder is not trained.")
if args.gradient_checkpointing:
unet.enable_gradient_checkpointing()
text_encoder.gradient_checkpointing_enable()
if not cache_latents:
vae.requires_grad_(False)
vae.eval()
vae.to(accelerator.device, dtype=weight_dtype)
# 学習に必要なクラスを準備する
print("prepare optimizer, data loader etc.")
if train_text_encoder:
trainable_params = itertools.chain(unet.parameters(), text_encoder.parameters())
else:
trainable_params = unet.parameters()
_, _, optimizer = train_util.get_optimizer(args, trainable_params)
# dataloaderを準備する
# DataLoaderのプロセス数:0はメインプロセスになる
n_workers = min(args.max_data_loader_n_workers, os.cpu_count() - 1) # cpu_count-1 ただし最大で指定された数まで
train_dataloader = torch.utils.data.DataLoader(
train_dataset_group,
batch_size=1,
shuffle=True,
collate_fn=collater,
num_workers=n_workers,
persistent_workers=args.persistent_data_loader_workers,
)
# 学習ステップ数を計算する
if args.max_train_epochs is not None:
args.max_train_steps = args.max_train_epochs * math.ceil(
len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps
)
print(f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}")
# データセット側にも学習ステップを送信
train_dataset_group.set_max_train_steps(args.max_train_steps)
if args.stop_text_encoder_training is None:
args.stop_text_encoder_training = args.max_train_steps + 1 # do not stop until end
# lr schedulerを用意する TODO gradient_accumulation_stepsの扱いが何かおかしいかもしれない。後で確認する
lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)
# 実験的機能:勾配も含めたfp16学習を行う モデル全体をfp16にする
if args.full_fp16:
assert (
args.mixed_precision == "fp16"
), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"
print("enable full fp16 training.")
unet.to(weight_dtype)
text_encoder.to(weight_dtype)
# acceleratorがなんかよろしくやってくれるらしい
if train_text_encoder:
unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
unet, text_encoder, optimizer, train_dataloader, lr_scheduler
)
else:
unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(unet, optimizer, train_dataloader, lr_scheduler)
# transform DDP after prepare
text_encoder, unet = train_util.transform_if_model_is_DDP(text_encoder, unet)
if not train_text_encoder:
text_encoder.to(accelerator.device, dtype=weight_dtype) # to avoid 'cpu' vs 'cuda' error
# 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする
if args.full_fp16:
train_util.patch_accelerator_for_fp16_training(accelerator)
# resumeする
train_util.resume_from_local_or_hf_if_specified(accelerator, args)
# epoch数を計算する
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0):
args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1
# 学習する
total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps
print("running training / 学習開始")
print(f" num train images * repeats / 学習画像の数×繰り返し回数: {train_dataset_group.num_train_images}")
print(f" num reg images / 正則化画像の数: {train_dataset_group.num_reg_images}")
print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}")
print(f" num epochs / epoch数: {num_train_epochs}")
print(f" batch size per device / バッチサイズ: {args.train_batch_size}")
print(f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}")
print(f" gradient ccumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}")
print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}")
progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps")
global_step = 0
noise_scheduler = DDPMScheduler(
beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False
)
prepare_scheduler_for_custom_training(noise_scheduler, accelerator.device)
if accelerator.is_main_process:
accelerator.init_trackers("dreambooth" if args.log_tracker_name is None else args.log_tracker_name)
if args.sample_first or args.sample_only:
# Do initial sample before starting training
train_tools.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizer,
text_encoder, unet, force_sample=True)
if args.sample_only:
return
loss_list = []
loss_total = 0.0
for epoch in range(num_train_epochs):
print(f"\nepoch {epoch+1}/{num_train_epochs}")
current_epoch.value = epoch + 1
# 指定したステップ数までText Encoderを学習する:epoch最初の状態
unet.train()
# train==True is required to enable gradient_checkpointing
if args.gradient_checkpointing or global_step < args.stop_text_encoder_training:
text_encoder.train()
for step, batch in enumerate(train_dataloader):
current_step.value = global_step
# 指定したステップ数でText Encoderの学習を止める
if global_step == args.stop_text_encoder_training:
print(f"stop text encoder training at step {global_step}")
if not args.gradient_checkpointing:
text_encoder.train(False)
text_encoder.requires_grad_(False)
with accelerator.accumulate(unet):
with torch.no_grad():
# latentに変換
if cache_latents:
latents = batch["latents"].to(accelerator.device)
else:
latents = vae.encode(batch["images"].to(dtype=weight_dtype)).latent_dist.sample()
latents = latents * 0.18215
b_size = latents.shape[0]
# Sample noise that we'll add to the latents
if args.train_noise_seed is not None:
torch.manual_seed(args.train_noise_seed)
torch.cuda.manual_seed(args.train_noise_seed)
# make same seed for each item in the batch by stacking them
single_noise = torch.randn_like(latents[0])
noise = torch.stack([single_noise for _ in range(b_size)])
noise = noise.to(latents.device)
elif args.seed_lock:
noise = train_tools.get_noise_from_latents(latents)
else:
noise = torch.randn_like(latents, device=latents.device)
if args.noise_offset:
noise = apply_noise_offset(latents, noise, args.noise_offset, args.adaptive_noise_scale)
elif args.multires_noise_iterations:
noise = pyramid_noise_like(noise, latents.device, args.multires_noise_iterations, args.multires_noise_discount)
# elif args.perlin_noise:
# noise = perlin_noise(noise, latents.device, args.perlin_noise) # only shape of noise is used currently
# Get the text embedding for conditioning
with torch.set_grad_enabled(global_step < args.stop_text_encoder_training):
if args.weighted_captions:
encoder_hidden_states = get_weighted_text_embeddings(
tokenizer,
text_encoder,
batch["captions"],
accelerator.device,
args.max_token_length // 75 if args.max_token_length else 1,
clip_skip=args.clip_skip,
)
else:
input_ids = batch["input_ids"].to(accelerator.device)
encoder_hidden_states = train_util.get_hidden_states(
args, input_ids, tokenizer, text_encoder, None if not args.full_fp16 else weight_dtype
)
# Sample a random timestep for each image
timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (b_size,), device=latents.device)
timesteps = timesteps.long()
# Add noise to the latents according to the noise magnitude at each timestep
# (this is the forward diffusion process)
noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
# Predict the noise residual
with accelerator.autocast():
noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample
if args.v_parameterization:
# v-parameterization training
target = noise_scheduler.get_velocity(latents, noise, timesteps)
else:
target = noise
loss = torch.nn.functional.mse_loss(noise_pred.float(), target.float(), reduction="none")
loss = loss.mean([1, 2, 3])
loss_weights = batch["loss_weights"] # 各sampleごとのweight
loss = loss * loss_weights
if args.min_snr_gamma:
loss = apply_snr_weight(loss, timesteps, noise_scheduler, args.min_snr_gamma)
if args.scale_v_pred_loss_like_noise_pred:
loss = scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler)
loss = loss.mean() # 平均なのでbatch_sizeで割る必要なし
accelerator.backward(loss)
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
if train_text_encoder:
params_to_clip = itertools.chain(unet.parameters(), text_encoder.parameters())
else:
params_to_clip = unet.parameters()
accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad(set_to_none=True)
# Checks if the accelerator has performed an optimization step behind the scenes
if accelerator.sync_gradients:
progress_bar.update(1)
global_step += 1
train_util.sample_images(
accelerator, args, None, global_step, accelerator.device, vae, tokenizer, text_encoder, unet
)
# 指定ステップごとにモデルを保存
if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0:
accelerator.wait_for_everyone()
if accelerator.is_main_process:
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
train_util.save_sd_model_on_epoch_end_or_stepwise(
args,
False,
accelerator,
src_path,
save_stable_diffusion_format,
use_safetensors,
save_dtype,
epoch,
num_train_epochs,
global_step,
unwrap_model(text_encoder),
unwrap_model(unet),
vae,
)
current_loss = loss.detach().item()
if args.logging_dir is not None:
logs = {"loss": current_loss, "lr": float(lr_scheduler.get_last_lr()[0])}
if args.optimizer_type.lower().startswith("DAdapt".lower()) or args.optimizer_type.lower() == "Prodigy".lower(): # tracking d*lr value
logs["lr/d*lr"] = (
lr_scheduler.optimizers[0].param_groups[0]["d"] * lr_scheduler.optimizers[0].param_groups[0]["lr"]
)
accelerator.log(logs, step=global_step)
if epoch == 0:
loss_list.append(current_loss)
else:
loss_total -= loss_list[step]
loss_list[step] = current_loss
loss_total += current_loss
avr_loss = loss_total / len(loss_list)
logs = {"loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}
progress_bar.set_postfix(**logs)
if global_step >= args.max_train_steps:
break
if args.logging_dir is not None:
logs = {"loss/epoch": loss_total / len(loss_list)}
accelerator.log(logs, step=epoch + 1)
accelerator.wait_for_everyone()
if args.save_every_n_epochs is not None:
if accelerator.is_main_process:
# checking for saving is in util
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
train_util.save_sd_model_on_epoch_end_or_stepwise(
args,
True,
accelerator,
src_path,
save_stable_diffusion_format,
use_safetensors,
save_dtype,
epoch,
num_train_epochs,
global_step,
unwrap_model(text_encoder),
unwrap_model(unet),
vae,
)
train_util.sample_images(accelerator, args, epoch + 1, global_step, accelerator.device, vae, tokenizer, text_encoder, unet)
is_main_process = accelerator.is_main_process
if is_main_process:
unet = unwrap_model(unet)
text_encoder = unwrap_model(text_encoder)
accelerator.end_training()
if args.save_state and is_main_process:
train_util.save_state_on_train_end(args, accelerator)
del accelerator # この後メモリを使うのでこれは消す
if is_main_process:
src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path
train_util.save_sd_model_on_train_end(
args, src_path, save_stable_diffusion_format, use_safetensors, save_dtype, epoch, global_step, text_encoder, unet, vae
)
print("model saved.")
def setup_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser()
train_util.add_sd_models_arguments(parser)
train_util.add_dataset_arguments(parser, True, False, True)
train_util.add_training_arguments(parser, True)
train_util.add_sd_saving_arguments(parser)
train_util.add_optimizer_arguments(parser)
config_util.add_config_arguments(parser)
custom_train_functions.add_custom_train_arguments(parser)
parser.add_argument(
"--no_token_padding",
action="store_true",
help="disable token padding (same as Diffuser's DreamBooth) / トークンのpaddingを無効にする(Diffusers版DreamBoothと同じ動作)",
)
parser.add_argument(
"--stop_text_encoder_training",
type=int,
default=None,
help="steps to stop text encoder training, -1 for no training / Text Encoderの学習を止めるステップ数、-1で最初から学習しない",
)
parser.add_argument(
"--sample_first",
action="store_true",
help="Sample first interval before training",
default=False
)
parser.add_argument(
"--name_replace",
type=str,
help="Replaces [name] in prompts. Used is sampling, training, and regs",
default=None
)
parser.add_argument(
"--train_noise_seed",
type=int,
help="Use custom seed for training noise",
default=None
)
parser.add_argument(
"--sample_only",
action="store_true",
help="Only generate samples. Used for generating training data with specific seeds to alter during training",
default=False
)
parser.add_argument(
"--seed_lock",
action="store_true",
help="Locks the seed to the latent images so the same latent will always have the same noise",
default=False
)
return parser
if __name__ == "__main__":
parser = setup_parser()
args = parser.parse_args()
args = train_util.read_config_from_file(args, parser)
train(args)

99
testing/compare_keys.py Normal file
View File

@@ -0,0 +1,99 @@
import argparse
import os
import torch
from diffusers.loaders import LoraLoaderMixin
from safetensors.torch import load_file
from collections import OrderedDict
import json
# this was just used to match the vae keys to the diffusers keys
# you probably wont need this. Unless they change them.... again... again
# on second thought, you probably will
device = torch.device('cpu')
dtype = torch.float32
parser = argparse.ArgumentParser()
# require at lease one config file
parser.add_argument(
'file_1',
nargs='+',
type=str,
help='Path to first safe tensor file'
)
parser.add_argument(
'file_2',
nargs='+',
type=str,
help='Path to second safe tensor file'
)
args = parser.parse_args()
find_matches = False
state_dict_file_1 = load_file(args.file_1[0])
state_dict_1_keys = list(state_dict_file_1.keys())
state_dict_file_2 = load_file(args.file_2[0])
state_dict_2_keys = list(state_dict_file_2.keys())
keys_in_both = []
keys_not_in_state_dict_2 = []
for key in state_dict_1_keys:
if key not in state_dict_2_keys:
keys_not_in_state_dict_2.append(key)
keys_not_in_state_dict_1 = []
for key in state_dict_2_keys:
if key not in state_dict_1_keys:
keys_not_in_state_dict_1.append(key)
keys_in_both = []
for key in state_dict_1_keys:
if key in state_dict_2_keys:
keys_in_both.append(key)
# sort them
keys_not_in_state_dict_2.sort()
keys_not_in_state_dict_1.sort()
keys_in_both.sort()
json_data = {
"both": keys_in_both,
"not_in_state_dict_2": keys_not_in_state_dict_2,
"not_in_state_dict_1": keys_not_in_state_dict_1
}
json_data = json.dumps(json_data, indent=4)
remaining_diffusers_values = OrderedDict()
for key in keys_not_in_state_dict_1:
remaining_diffusers_values[key] = state_dict_file_2[key]
# print(remaining_diffusers_values.keys())
remaining_ldm_values = OrderedDict()
for key in keys_not_in_state_dict_2:
remaining_ldm_values[key] = state_dict_file_1[key]
# print(json_data)
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
json_save_path = os.path.join(project_root, 'config', 'keys.json')
json_matched_save_path = os.path.join(project_root, 'config', 'matched.json')
json_duped_save_path = os.path.join(project_root, 'config', 'duped.json')
state_dict_1_filename = os.path.basename(args.file_1[0])
state_dict_2_filename = os.path.basename(args.file_2[0])
# save key names for each in own file
with open(os.path.join(project_root, 'config', f'{state_dict_1_filename}.json'), 'w') as f:
f.write(json.dumps(state_dict_1_keys, indent=4))
with open(os.path.join(project_root, 'config', f'{state_dict_2_filename}.json'), 'w') as f:
f.write(json.dumps(state_dict_2_keys, indent=4))
with open(json_save_path, 'w') as f:
f.write(json_data)

View File

@@ -0,0 +1,130 @@
from collections import OrderedDict
import torch
from safetensors.torch import load_file
import argparse
import os
import json
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
keymap_path = os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', 'stable_diffusion_sdxl.json')
# load keymap
with open(keymap_path, 'r') as f:
keymap = json.load(f)
lora_keymap = OrderedDict()
# convert keymap to lora key naming
for ldm_key, diffusers_key in keymap['ldm_diffusers_keymap'].items():
if ldm_key.endswith('.bias') or diffusers_key.endswith('.bias'):
# skip it
continue
# sdxl has same te for locon with kohya and ours
if ldm_key.startswith('conditioner'):
#skip it
continue
# ignore vae
if ldm_key.startswith('first_stage_model'):
continue
ldm_key = ldm_key.replace('model.diffusion_model.', 'lora_unet_')
ldm_key = ldm_key.replace('.weight', '')
ldm_key = ldm_key.replace('.', '_')
diffusers_key = diffusers_key.replace('unet_', 'lora_unet_')
diffusers_key = diffusers_key.replace('.weight', '')
diffusers_key = diffusers_key.replace('.', '_')
lora_keymap[f"{ldm_key}.alpha"] = f"{diffusers_key}.alpha"
lora_keymap[f"{ldm_key}.lora_down.weight"] = f"{diffusers_key}.lora_down.weight"
lora_keymap[f"{ldm_key}.lora_up.weight"] = f"{diffusers_key}.lora_up.weight"
parser = argparse.ArgumentParser()
parser.add_argument("input", help="input file")
parser.add_argument("input2", help="input2 file")
args = parser.parse_args()
# name = args.name
# if args.sdxl:
# name += '_sdxl'
# elif args.sd2:
# name += '_sd2'
# else:
# name += '_sd1'
name = 'stable_diffusion_locon_sdxl'
locon_save = load_file(args.input)
our_save = load_file(args.input2)
our_extra_keys = list(set(our_save.keys()) - set(locon_save.keys()))
locon_extra_keys = list(set(locon_save.keys()) - set(our_save.keys()))
print(f"we have {len(our_extra_keys)} extra keys")
print(f"locon has {len(locon_extra_keys)} extra keys")
save_dtype = torch.float16
print(f"our extra keys: {our_extra_keys}")
print(f"locon extra keys: {locon_extra_keys}")
def export_state_dict(our_save):
converted_state_dict = OrderedDict()
for key, value in our_save.items():
# test encoders share keys for some reason
if key.startswith('lora_te'):
converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
else:
converted_key = key
for ldm_key, diffusers_key in lora_keymap.items():
if converted_key == diffusers_key:
converted_key = ldm_key
converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
return converted_state_dict
def import_state_dict(loaded_state_dict):
converted_state_dict = OrderedDict()
for key, value in loaded_state_dict.items():
if key.startswith('lora_te'):
converted_state_dict[key] = value.detach().to('cpu', dtype=save_dtype)
else:
converted_key = key
for ldm_key, diffusers_key in lora_keymap.items():
if converted_key == ldm_key:
converted_key = diffusers_key
converted_state_dict[converted_key] = value.detach().to('cpu', dtype=save_dtype)
return converted_state_dict
# check it again
converted_state_dict = export_state_dict(our_save)
converted_extra_keys = list(set(converted_state_dict.keys()) - set(locon_save.keys()))
locon_extra_keys = list(set(locon_save.keys()) - set(converted_state_dict.keys()))
print(f"we have {len(converted_extra_keys)} extra keys")
print(f"locon has {len(locon_extra_keys)} extra keys")
print(f"our extra keys: {converted_extra_keys}")
# convert back
cycle_state_dict = import_state_dict(converted_state_dict)
cycle_extra_keys = list(set(cycle_state_dict.keys()) - set(our_save.keys()))
our_extra_keys = list(set(our_save.keys()) - set(cycle_state_dict.keys()))
print(f"we have {len(our_extra_keys)} extra keys")
print(f"cycle has {len(cycle_extra_keys)} extra keys")
# save keymap
to_save = OrderedDict()
to_save['ldm_diffusers_keymap'] = lora_keymap
with open(os.path.join(PROJECT_ROOT, 'toolkit', 'keymaps', f'{name}.json'), 'w') as f:
json.dump(to_save, f, indent=4)

View File

@@ -0,0 +1,479 @@
import argparse
import gc
import os
import re
import os
# add project root to sys path
import sys
from diffusers import DiffusionPipeline, StableDiffusionXLPipeline
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from diffusers.loaders import LoraLoaderMixin
from safetensors.torch import load_file, save_file
from collections import OrderedDict
import json
from tqdm import tqdm
from toolkit.config_modules import ModelConfig
from toolkit.stable_diffusion_model import StableDiffusion
KEYMAPS_FOLDER = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), 'toolkit', 'keymaps')
device = torch.device('cpu')
dtype = torch.float32
def flush():
torch.cuda.empty_cache()
gc.collect()
def get_reduced_shape(shape_tuple):
# iterate though shape anr remove 1s
new_shape = []
for dim in shape_tuple:
if dim != 1:
new_shape.append(dim)
return tuple(new_shape)
parser = argparse.ArgumentParser()
# require at lease one config file
parser.add_argument(
'file_1',
nargs='+',
type=str,
help='Path to first safe tensor file'
)
parser.add_argument('--name', type=str, default='stable_diffusion', help='name for mapping to make')
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()
file_path = args.file_1[0]
find_matches = False
print(f'Loading diffusers model')
ignore_ldm_begins_with = []
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"
if not args.refiner:
diffusers_model_config = ModelConfig(
name_or_path=diffusers_file_path,
is_xl=args.sdxl,
is_v2=args.sd2,
is_ssd=args.ssd,
is_vega=args.vega,
dtype=dtype,
)
diffusers_sd = StableDiffusion(
model_config=diffusers_model_config,
device=device,
dtype=dtype,
)
diffusers_sd.load_model()
# delete things we dont need
del diffusers_sd.tokenizer
flush()
print(f'Loading ldm model')
diffusers_state_dict = diffusers_sd.state_dict()
else:
# refiner wont work directly with stable diffusion
# so we need to load the model and then load the state dict
diffusers_pipeline = StableDiffusionXLPipeline.from_single_file(
diffusers_file_path,
torch_dtype=torch.float16,
use_safetensors=True,
variant="fp16",
).to(device)
# diffusers_pipeline = StableDiffusionXLPipeline.from_single_file(
# file_path,
# torch_dtype=torch.float16,
# use_safetensors=True,
# variant="fp16",
# ).to(device)
SD_PREFIX_VAE = "vae"
SD_PREFIX_UNET = "unet"
SD_PREFIX_REFINER_UNET = "refiner_unet"
SD_PREFIX_TEXT_ENCODER = "te"
SD_PREFIX_TEXT_ENCODER1 = "te0"
SD_PREFIX_TEXT_ENCODER2 = "te1"
diffusers_state_dict = OrderedDict()
for k, v in diffusers_pipeline.vae.state_dict().items():
new_key = k if k.startswith(f"{SD_PREFIX_VAE}") else f"{SD_PREFIX_VAE}_{k}"
diffusers_state_dict[new_key] = v
for k, v in diffusers_pipeline.text_encoder_2.state_dict().items():
new_key = k if k.startswith(f"{SD_PREFIX_TEXT_ENCODER2}_") else f"{SD_PREFIX_TEXT_ENCODER2}_{k}"
diffusers_state_dict[new_key] = v
for k, v in diffusers_pipeline.unet.state_dict().items():
new_key = k if k.startswith(f"{SD_PREFIX_UNET}_") else f"{SD_PREFIX_UNET}_{k}"
diffusers_state_dict[new_key] = v
# add ignore ones as we are only going to focus on unet and copy the rest
# ignore_ldm_begins_with = ["conditioner.", "first_stage_model."]
diffusers_dict_keys = list(diffusers_state_dict.keys())
ldm_state_dict = load_file(file_path)
ldm_dict_keys = list(ldm_state_dict.keys())
ldm_diffusers_keymap = OrderedDict()
ldm_diffusers_shape_map = OrderedDict()
ldm_operator_map = OrderedDict()
diffusers_operator_map = OrderedDict()
total_keys = len(ldm_dict_keys)
matched_ldm_keys = []
matched_diffusers_keys = []
error_margin = 1e-8
tmp_merge_key = "TMP___MERGE"
te_suffix = ''
proj_pattern_weight = None
proj_pattern_bias = None
text_proj_layer = None
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"
proj_pattern_bias = r"conditioner\.embedders\.1\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
text_proj_layer = "conditioner.embedders.1.model.text_projection"
if args.refiner:
te_suffix = '1'
ldm_res_block_prefix = "conditioner.embedders.0.model.transformer.resblocks"
proj_pattern_weight = r"conditioner\.embedders\.0\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
proj_pattern_bias = r"conditioner\.embedders\.0\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_bias"
text_proj_layer = "conditioner.embedders.0.model.text_projection"
if args.sd2:
te_suffix = ''
ldm_res_block_prefix = "cond_stage_model.model.transformer.resblocks"
proj_pattern_weight = r"cond_stage_model\.model\.transformer\.resblocks\.(\d+)\.attn\.in_proj_weight"
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 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])
else:
d_model = 1024
# do pre known merging
for ldm_key in ldm_dict_keys:
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"],
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight"],
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight"],
], dim=0)
# add to matched so we dont check them
matched_diffusers_keys.append(
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight")
matched_diffusers_keys.append(
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight")
matched_diffusers_keys.append(
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight")
# make diffusers convertable_dict
diffusers_state_dict[
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.{tmp_merge_key}.weight"] = new_val
# add operator
ldm_operator_map[ldm_key] = {
"cat": [
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight",
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight",
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight",
],
}
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:, :]
# add diffusers operators
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.weight"] = {
"slice": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
f"0:{d_model}, :"
]
}
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.weight"] = {
"slice": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
f"{d_model}:{d_model * 2}, :"
]
}
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.weight"] = {
"slice": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_weight",
f"{d_model * 2}:, :"
]
}
match = re.match(proj_pattern_bias, ldm_key)
if match:
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.bias"],
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias"],
diffusers_state_dict[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias"],
], dim=0)
# add to matched so we dont check them
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias")
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias")
matched_diffusers_keys.append(f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias")
# make diffusers convertable_dict
diffusers_state_dict[
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.{tmp_merge_key}.bias"] = new_val
# add operator
ldm_operator_map[ldm_key] = {
"cat": [
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.q_proj.bias",
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias",
f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias",
],
}
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": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
f"0:{d_model}, :"
]
}
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.k_proj.bias"] = {
"slice": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
f"{d_model}:{d_model * 2}, :"
]
}
diffusers_operator_map[f"te{te_suffix}_text_model.encoder.layers.{number}.self_attn.v_proj.bias"] = {
"slice": [
f"{ldm_res_block_prefix}.{number}.attn.in_proj_bias",
f"{d_model * 2}:, :"
]
}
except Exception as e:
print(f"Error on key {ldm_key}")
print(e)
# update keys
diffusers_dict_keys = list(diffusers_state_dict.keys())
pbar = tqdm(ldm_dict_keys, desc='Matching ldm-diffusers keys', total=total_keys)
# run through all weights and check mse between them to find matches
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)
# That was easy. Same key
# if ldm_key == diffusers_key:
# ldm_diffusers_keymap[ldm_key] = diffusers_key
# matched_ldm_keys.append(ldm_key)
# matched_diffusers_keys.append(diffusers_key)
# break
# if we already have this key mapped, skip it
if diffusers_key in matched_diffusers_keys:
continue
# if reduced shapes do not match skip it
if ldm_reduced_shape_tuple != diffusers_reduced_shape_tuple:
continue
ldm_weight = ldm_state_dict[ldm_key]
did_reduce_ldm = False
diffusers_weight = diffusers_state_dict[diffusers_key]
did_reduce_diffusers = False
# reduce the shapes to match if they are not the same
if ldm_shape_tuple != ldm_reduced_shape_tuple:
ldm_weight = ldm_weight.view(ldm_reduced_shape_tuple)
did_reduce_ldm = True
if diffusers_shape_tuple != diffusers_reduced_shape_tuple:
diffusers_weight = diffusers_weight.view(diffusers_reduced_shape_tuple)
did_reduce_diffusers = True
# check to see if they match within a margin of error
mse = torch.nn.functional.mse_loss(ldm_weight.float(), diffusers_weight.float())
if mse < error_margin:
ldm_diffusers_keymap[ldm_key] = diffusers_key
matched_ldm_keys.append(ldm_key)
matched_diffusers_keys.append(diffusers_key)
if did_reduce_ldm or did_reduce_diffusers:
ldm_diffusers_shape_map[ldm_key] = (ldm_shape_tuple, diffusers_shape_tuple)
if did_reduce_ldm:
del ldm_weight
if did_reduce_diffusers:
del diffusers_weight
flush()
break
pbar.update(1)
pbar.close()
name = args.name
if args.sdxl:
name += '_sdxl'
elif args.ssd:
name += '_ssd'
elif args.vega:
name += '_vega'
elif args.refiner:
name += '_refiner'
elif args.sd2:
name += '_sd2'
else:
name += '_sd1'
# if len(matched_ldm_keys) != len(matched_diffusers_keys):
unmatched_ldm_keys = [x for x in ldm_dict_keys if x not in matched_ldm_keys]
unmatched_diffusers_keys = [x for x in diffusers_dict_keys if x not in matched_diffusers_keys]
# has unmatched keys
has_unmatched_keys = len(unmatched_ldm_keys) > 0 or len(unmatched_diffusers_keys) > 0
def get_slices_from_string(s: str) -> tuple:
slice_strings = s.split(',')
slices = [eval(f"slice({component.strip()})") for component in slice_strings]
return tuple(slices)
if has_unmatched_keys:
print(
f"Found {len(unmatched_ldm_keys)} unmatched ldm keys and {len(unmatched_diffusers_keys)} unmatched diffusers keys")
unmatched_obj = OrderedDict()
unmatched_obj['ldm'] = OrderedDict()
unmatched_obj['diffusers'] = OrderedDict()
print(f"Gathering info on unmatched keys")
for key in tqdm(unmatched_ldm_keys, desc='Unmatched LDM keys'):
# get min, max, mean, std
weight = ldm_state_dict[key]
weight_min = weight.min().item()
weight_max = weight.max().item()
unmatched_obj['ldm'][key] = {
'shape': weight.shape,
"min": weight_min,
"max": weight_max,
}
del weight
flush()
for key in tqdm(unmatched_diffusers_keys, desc='Unmatched Diffusers keys'):
# get min, max, mean, std
weight = diffusers_state_dict[key]
weight_min = weight.min().item()
weight_max = weight.max().item()
unmatched_obj['diffusers'][key] = {
"shape": weight.shape,
"min": weight_min,
"max": weight_max,
}
del weight
flush()
unmatched_path = os.path.join(KEYMAPS_FOLDER, f'{name}_unmatched.json')
with open(unmatched_path, 'w') as f:
f.write(json.dumps(unmatched_obj, indent=4))
print(f'Saved unmatched keys to {unmatched_path}')
# save ldm remainders
remaining_ldm_values = OrderedDict()
for key in unmatched_ldm_keys:
remaining_ldm_values[key] = ldm_state_dict[key].detach().to('cpu', torch.float16)
save_file(remaining_ldm_values, os.path.join(KEYMAPS_FOLDER, f'{name}_ldm_base.safetensors'))
print(f'Saved remaining ldm values to {os.path.join(KEYMAPS_FOLDER, f"{name}_ldm_base.safetensors")}')
# do cleanup of some left overs and bugs
to_remove = []
for ldm_key, diffusers_key in ldm_diffusers_keymap.items():
# get rid of tmp merge keys used to slicing
if tmp_merge_key in diffusers_key or tmp_merge_key in ldm_key:
to_remove.append(ldm_key)
for key in to_remove:
del ldm_diffusers_keymap[key]
to_remove = []
# remove identical shape mappings. Not sure why they exist but they do
for ldm_key, shape_list in ldm_diffusers_shape_map.items():
# remove identical shape mappings. Not sure why they exist but they do
# convert to json string to make it easier to compare
ldm_shape = json.dumps(shape_list[0])
diffusers_shape = json.dumps(shape_list[1])
if ldm_shape == diffusers_shape:
to_remove.append(ldm_key)
for key in to_remove:
del ldm_diffusers_shape_map[key]
dest_path = os.path.join(KEYMAPS_FOLDER, f'{name}.json')
save_obj = OrderedDict()
save_obj["ldm_diffusers_keymap"] = ldm_diffusers_keymap
save_obj["ldm_diffusers_shape_map"] = ldm_diffusers_shape_map
save_obj["ldm_diffusers_operator_map"] = ldm_operator_map
save_obj["diffusers_ldm_operator_map"] = diffusers_operator_map
with open(dest_path, 'w') as f:
f.write(json.dumps(save_obj, indent=4))
print(f'Saved keymap to {dest_path}')

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

@@ -0,0 +1,128 @@
import time
import numpy as np
import torch
from torch.utils.data import DataLoader
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
import torchvision.transforms.functional
from toolkit.image_utils import show_img, show_tensors
sys.path.append(SD_SCRIPTS_ROOT)
from library.model_util import load_vae
from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
from toolkit.data_loader import AiToolkitDataset, get_dataloader_from_datasets, \
trigger_dataloader_setup_epoch
from toolkit.config_modules import DatasetConfig
import argparse
from tqdm import tqdm
parser = argparse.ArgumentParser()
parser.add_argument('dataset_folder', type=str, default='input')
parser.add_argument('--epochs', type=int, default=1)
args = parser.parse_args()
dataset_folder = args.dataset_folder
resolution = 1024
bucket_tolerance = 64
batch_size = 1
clip_processor = CLIPImageProcessor.from_pretrained("openai/clip-vit-base-patch16")
class FakeAdapter:
def __init__(self):
self.clip_image_processor = clip_processor
## make fake sd
class FakeSD:
def __init__(self):
self.adapter = FakeAdapter()
dataset_config = DatasetConfig(
dataset_path=dataset_folder,
# clip_image_path=dataset_folder,
# square_crop=True,
resolution=resolution,
# caption_ext='json',
default_caption='default',
# clip_image_path='/mnt/Datasets2/regs/yetibear_xl_v14/random_aspect/',
buckets=True,
bucket_tolerance=bucket_tolerance,
# 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, sd=FakeSD())
# run through an epoch ang check sizes
dataloader_iterator = iter(dataloader)
for epoch in range(args.epochs):
for batch in tqdm(dataloader):
batch: 'DataLoaderBatchDTO'
img_batch = batch.tensor
batch_size, channels, height, width = img_batch.shape
# img_batch = color_block_imgs(img_batch, neg1_1=True)
# 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 = img_batch
# big_img = big_img.clamp(-1, 1)
show_tensors(big_img)
# convert to image
# img = transforms.ToPILImage()(big_img)
#
# show_img(img)
time.sleep(0.2)
# if not last epoch
if epoch < args.epochs - 1:
trigger_dataloader_setup_epoch(dataloader)
cv2.destroyAllWindows()
print('done')

View File

@@ -0,0 +1,172 @@
import argparse
import os
# add project root to sys path
import sys
from tqdm import tqdm
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
from diffusers.loaders import LoraLoaderMixin
from safetensors.torch import load_file
from collections import OrderedDict
import json
from toolkit.config_modules import ModelConfig
from toolkit.paths import KEYMAPS_ROOT
from toolkit.saving import convert_state_dict_to_ldm_with_mapping, get_ldm_state_dict_from_diffusers
from toolkit.stable_diffusion_model import StableDiffusion
# this was just used to match the vae keys to the diffusers keys
# you probably wont need this. Unless they change them.... again... again
# on second thought, you probably will
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
device = torch.device('cpu')
dtype = torch.float32
parser = argparse.ArgumentParser()
# require at lease one config file
parser.add_argument(
'file_1',
nargs='+',
type=str,
help='Path an LDM model'
)
parser.add_argument(
'--is_xl',
action='store_true',
help='Is the model an XL model'
)
parser.add_argument(
'--is_v2',
action='store_true',
help='Is the model a v2 model'
)
args = parser.parse_args()
find_matches = False
print("Loading model")
state_dict_file_1 = load_file(args.file_1[0])
state_dict_1_keys = list(state_dict_file_1.keys())
print("Loading model into diffusers format")
model_config = ModelConfig(
name_or_path=args.file_1[0],
is_xl=args.is_xl
)
sd = StableDiffusion(
model_config=model_config,
device=device,
)
sd.load_model()
# load our base
base_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl_ldm_base.safetensors')
mapping_path = os.path.join(KEYMAPS_ROOT, 'stable_diffusion_sdxl.json')
print("Converting model back to LDM")
version_string = '1'
if args.is_v2:
version_string = '2'
if args.is_xl:
version_string = 'sdxl'
# convert the state dict
state_dict_file_2 = get_ldm_state_dict_from_diffusers(
sd.state_dict(),
version_string,
device='cpu',
dtype=dtype
)
# state_dict_file_2 = load_file(args.file_2[0])
state_dict_2_keys = list(state_dict_file_2.keys())
keys_in_both = []
keys_not_in_state_dict_2 = []
for key in state_dict_1_keys:
if key not in state_dict_2_keys:
keys_not_in_state_dict_2.append(key)
keys_not_in_state_dict_1 = []
for key in state_dict_2_keys:
if key not in state_dict_1_keys:
keys_not_in_state_dict_1.append(key)
keys_in_both = []
for key in state_dict_1_keys:
if key in state_dict_2_keys:
keys_in_both.append(key)
# sort them
keys_not_in_state_dict_2.sort()
keys_not_in_state_dict_1.sort()
keys_in_both.sort()
if len(keys_not_in_state_dict_2) == 0 and len(keys_not_in_state_dict_1) == 0:
print("All keys match!")
print("Checking values...")
mismatch_keys = []
loss = torch.nn.MSELoss()
tolerance = 1e-6
for key in tqdm(keys_in_both):
if loss(state_dict_file_1[key], state_dict_file_2[key]) > tolerance:
print(f"Values for key {key} don't match!")
print(f"Loss: {loss(state_dict_file_1[key], state_dict_file_2[key])}")
mismatch_keys.append(key)
if len(mismatch_keys) == 0:
print("All values match!")
else:
print("Some valued font match!")
print(mismatch_keys)
mismatched_path = os.path.join(project_root, 'config', 'mismatch.json')
with open(mismatched_path, 'w') as f:
f.write(json.dumps(mismatch_keys, indent=4))
exit(0)
else:
print("Keys don't match!, generating info...")
json_data = {
"both": keys_in_both,
"not_in_state_dict_2": keys_not_in_state_dict_2,
"not_in_state_dict_1": keys_not_in_state_dict_1
}
json_data = json.dumps(json_data, indent=4)
remaining_diffusers_values = OrderedDict()
for key in keys_not_in_state_dict_1:
remaining_diffusers_values[key] = state_dict_file_2[key]
# print(remaining_diffusers_values.keys())
remaining_ldm_values = OrderedDict()
for key in keys_not_in_state_dict_2:
remaining_ldm_values[key] = state_dict_file_1[key]
# print(json_data)
json_save_path = os.path.join(project_root, 'config', 'keys.json')
json_matched_save_path = os.path.join(project_root, 'config', 'matched.json')
json_duped_save_path = os.path.join(project_root, 'config', 'duped.json')
state_dict_1_filename = os.path.basename(args.file_1[0])
# state_dict_2_filename = os.path.basename(args.file_2[0])
# save key names for each in own file
with open(os.path.join(project_root, 'config', f'{state_dict_1_filename}.json'), 'w') as f:
f.write(json.dumps(state_dict_1_keys, indent=4))
with open(os.path.join(project_root, 'config', f'{state_dict_1_filename}_loop.json'), 'w') as f:
f.write(json.dumps(state_dict_2_keys, indent=4))
with open(json_save_path, 'w') as f:
f.write(json_data)

113
testing/test_vae.py Normal file
View File

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

55
toolkit/assistant_lora.py Normal file
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

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