321 Commits
Author SHA1 Message Date
mayukhdeb fb8ee98a1c penalize mean 2024-08-02 04:31:01 -07:00
mayukhdeb 7ff02a35ce smaller daam loss scale 2024-08-02 04:13:23 -07:00
mayukhdeb 21200b41b9 naive distribution loss, did not work :( 2024-08-02 03:17:32 -07:00
mayukhdeb c21605bb95 also plot token as string on title 2024-08-01 00:49:47 -07:00
mayukhdeb 4880e3f096 vis token wise daam maps 2024-08-01 00:39:49 -07:00
mayukhdeb 35e76fdef0 wip reproduce daam 2024-07-30 23:01:14 -07:00
mayukhdeb 871e425f3e obtain single layer heatmap 2024-07-30 22:50:56 -07:00
mayukhdeb e6534f8d1d baby steps 2024-07-30 22:50:22 -07:00
mayukhdeb b5ad0b659e watch daam loss and plot norms into heatmaps dir 2024-07-29 10:46:11 -07:00
mayukhdeb d6a205cc7b ahijack AttnProcessor2_0 2024-07-29 10:12:43 -07:00
mayukhdeb f21efd7840 train TI on 2 tokens 2024-07-29 10:11:15 -07:00
mayukhdeb 29855d94ba ignore notebook checkpoints folder 2024-07-22 12:03:25 -07:00
xander 006a3b9750 add ti config 2024-07-22 20:53:49 +02:00
aiXander b166f7c274 merge 2024-07-22 11:47:38 -07:00
aiXander f6cfd6d0bc fix typo 2024-07-22 11:46:33 -07:00
xander 6a4fa1ef8d add progressbar to node 2024-07-18 21:04:43 +02:00
xander 3cd17086f3 trainer comfyui v1 2024-07-18 05:50:04 +02:00
xander 75137a7539 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:07:07 +02:00
xander ee91f8c0c7 update naming conventions 2024-07-18 04:07:00 +02:00
xander 598c4204be add test config 2024-07-18 04:05:10 +02:00
xander 8da6c91af8 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:00:30 +02:00
xander 01523350e1 cleanup save dir 2024-07-18 04:00:21 +02:00
xander c891203365 update gitignore 2024-07-18 04:00:08 +02:00
xander 85599dc725 minor changes 2024-07-18 03:43:43 +02:00
xander af16d1edcc update configs 2024-07-18 03:26:51 +02:00
xander b3db4ff3bb Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-07-18 03:16:43 +02:00
xander 2a0761ae75 use gpt-4o 2024-07-18 03:16:41 +02:00
xander d8ee9ee522 use gpt-4o 2024-07-18 03:14:31 +02:00
xander 4f084d251c updates for comfyui node 2024-07-18 03:12:11 +02:00
aiXander fb521e9dbd add print flush 2024-07-16 04:39:29 -07:00
aiXander f937106f6c push changes 2024-07-16 04:04:34 -07:00
Gene Kogan 5f8fe5c4c8 Update requirements.txt 2024-07-16 03:44:59 -07:00
Gene Kogan 2f5eaeba7a Update cog.yaml 2024-07-16 03:44:35 -07:00
aiXander 7d8c845765 update yaml and print deps in config 2024-07-12 05:16:06 -07:00
aiXander 340336b53d print config pre training start 2024-07-11 06:11:33 -07:00
xander 5abed1487f setup lora_scale automation for validation grid 2024-07-11 13:36:11 +02:00
xander b5cf857b1f push small changes 2024-07-09 18:33:52 +02:00
xander e1813c75d1 updates to avoid OOM when plotting hist 2024-07-06 22:18:47 +02:00
xander 4fde1b7dd6 small tweaks 2024-07-06 18:45:38 +02:00
xander 892cf61c52 add disable_ti flag 2024-07-06 18:43:03 +02:00
xander fdd531d3c5 fix full finetuning and add 8bitadam 2024-07-06 17:42:49 +02:00
xander ef6635fafb update train cmd 2024-07-05 17:11:06 +02:00
xander faf864ce15 tiny tweak to learning rates for SDXL 2024-07-05 17:07:33 +02:00
xander b02221b23c cleanup before sd3 integration 2024-06-13 14:24:32 +02:00
xander d8b3175b57 cleanup before sd3 integration 2024-06-13 13:57:20 +02:00
aiXander a80c934679 add name 2024-06-05 05:47:26 -07:00
aiXander 813e95592d initial commit of ComfyUI node 2024-06-05 03:30:18 -07:00
xander 41ddfbcca6 add resource monitoring 2024-05-29 23:33:42 +02:00
xander 8cd935251d add strip urls 2024-05-28 23:11:41 +02:00
xander 0ed83dc656 update readme 2024-05-28 17:33:13 +02:00
Xander Steenbrugge 3ebbcf726b Update README.md 2024-05-28 17:28:31 +02:00
Xander Steenbrugge 3883adb697 Update README.md 2024-05-28 17:28:15 +02:00
xander 2033b13370 update readme 2024-05-28 17:26:46 +02:00
xander 3dabbfd55c add description for modes 2024-05-28 17:11:11 +02:00
xander 0d8f5ca929 update readme 2024-05-28 17:09:30 +02:00
xander 128adb99ab add imgs 2024-05-28 17:07:55 +02:00
xander 46ce69a29a update readme 2024-05-28 17:04:17 +02:00
xander 135959ba0e update readme 2024-05-28 16:56:56 +02:00
xander c845c76b15 update readme 2024-05-28 16:55:49 +02:00
xander f6f50e5004 merge 2024-05-28 16:51:43 +02:00
xander e9b5ca0dcc small tweaks 2024-05-28 16:51:02 +02:00
aiXander eb144d5212 increase starting lr for sd15 2024-05-16 16:46:08 -07:00
aiXander 3f7ba82a25 tiny fixes 2024-05-16 16:02:46 -07:00
aiXander 6e32be4dfd bugfix 2024-05-16 14:31:57 -07:00
aiXander 3e17bed208 update workflows 2024-05-16 14:03:14 -07:00
aiXander 9102e8e5e5 render with active pipe when not in debug mode 2024-05-16 09:28:22 -07:00
aiXander 03c6175007 fallback to blip when there are many imgs to caption 2024-05-16 09:04:20 -07:00
xander 13ee2ad6b5 plot categorical labels 2024-05-16 16:38:52 +02:00
aiXander e665169ce8 use gpt4-o 2024-05-16 07:37:02 -07:00
aiXander de99482825 fix predict cog bugs 2024-05-15 17:43:17 -07:00
aiXander 0d0ad2a4d0 update base model 2024-05-15 16:54:49 -07:00
aiXander 33b4ac7cc2 update default settings 2024-05-15 16:42:55 -07:00
aiXander 3f7298633d update cog run settings 2024-05-15 14:59:48 -07:00
aiXander fec9cbfadd fix cog predict 2024-05-14 15:11:07 -07:00
aiXander 06b3b7b459 print debug state 2024-05-14 07:07:13 -07:00
aiXander 70a645ac79 fix bugs 2024-05-14 07:04:25 -07:00
xander fa76a90c39 add workflows 2024-05-14 15:25:35 +02:00
xander 8b3f9a930f rename workflows 2024-05-14 15:24:33 +02:00
xander cfde9e9180 tiny param tweak 2024-05-14 15:23:26 +02:00
xander 4389ae75c8 force quit training after max steps 2024-05-14 15:21:48 +02:00
xander d66863b4cb more cleanup 2024-05-14 15:12:49 +02:00
xander 110ec0df30 fix lora_scale inference validation samples 2024-05-14 14:34:57 +02:00
xander b51727219e cleanup and try to fix last inference step 2024-05-14 11:42:57 +02:00
xander b8fe89e983 fix typo 2024-05-02 17:11:54 +02:00
aiXander e47fcf429e update predict.py 2024-05-02 07:59:13 -07:00
xander fc63399317 updates 2024-05-02 16:53:49 +02:00
xander 0bd26033cd add conv2 to lora layers 2024-04-29 20:12:08 +02:00
xander 7d78a3d76d updates 2024-04-29 19:37:45 +02:00
xander 45e6f50c24 small updates 2024-04-29 12:10:17 +02:00
xander 39ee9c9630 set defaults 2024-04-27 16:47:35 +02:00
xander 24fcc2e791 update defaults 2024-04-26 15:59:18 +02:00
xander 8da2f58c20 everything is working 2024-04-26 14:52:32 +02:00
xander 7d469d843c cleanup and make all this work also for sd15 2024-04-26 14:38:25 +02:00
xander 2e5912f629 improve token regularization 2024-04-26 13:32:01 +02:00
xander 37bcfe61c9 hyperparam tweaks 2024-04-26 08:45:29 +02:00
xander af6c4676d5 bugfixes 2024-04-25 21:41:57 +02:00
xander 9ae9bc19d2 mask improvements 2024-04-25 21:21:14 +02:00
xander 7218c7e8c2 more tok reg updates 2024-04-24 17:51:42 +02:00
xander c60f3646fa update tok_reg loss 2024-04-24 16:43:22 +02:00
xander e1b05b1b99 update loss plots, reduce tok_reg value 2024-04-24 14:42:05 +02:00
xander 0a72748da6 update preprocess to work better for 1 img 2024-04-23 16:21:49 +02:00
xander f919e3e7a1 setup experiment run 2024-04-23 04:39:35 +02:00
aiXander d4923ab6d2 tiny typo fix 2024-04-22 19:39:16 -07:00
aiXander d613eff061 more tweaks and cleanup 2024-04-22 19:24:28 -07:00
aiXander b48419247a refactor token regularization 2024-04-22 19:08:36 -07:00
aiXander 4d9f8c2e67 more tweaks 2024-04-22 12:22:20 -07:00
aiXander 34e4bc227c ready for pushing new concept trainer 2024-04-22 12:00:40 -07:00
aiXander c85c45b688 add a few more params to cog input 2024-04-22 11:41:37 -07:00
aiXander d4dbb52f9b add adapter_config.json save 2024-04-22 11:32:46 -07:00
aiXander 784bfbf14f prep for cog push 2024-04-22 11:27:04 -07:00
aiXander 327339ebba prep for cog push 2024-04-22 08:41:00 -07:00
aiXander f2d70b64da add dockerignore 2024-04-22 05:48:19 -07:00
xander 41b7895bff add unet_prodigy_growth_factor 2024-04-22 14:36:04 +02:00
mayukhdeb 5b4020d448 save sd15 models as they're supposed to be saved 2024-04-20 06:21:24 -07:00
aiXander 2633c06a87 update gpt4 model 2024-04-19 07:36:51 -07:00
mayukhdeb a33a33295e never load from disk during training 2024-04-19 05:36:24 -07:00
aiXander 3e44dff1da remove pipe before loading from disk 2024-04-19 05:19:49 -07:00
aiXander b2080da7c8 cleanup checkpoint 2024-04-19 05:08:47 -07:00
aiXander 13733a4f19 Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-04-19 04:58:09 -07:00
aiXander ccaeab8e71 freeze_ti param, save txt-encoder loras into savetensors file 2024-04-19 04:57:59 -07:00
mayukhdeb b568c0c283 load ;checkpoint from disk once training is complete 2024-04-19 04:55:02 -07:00
mayukhdeb c071905658 black formatting 2024-04-19 02:31:43 -07:00
aiXander d1d8eb125e fix 2024-04-19 01:41:28 -07:00
aiXander ea67904a9a simplify 2024-04-19 01:39:43 -07:00
aiXander 909e1c0b10 update defaults 2024-04-18 09:23:47 -07:00
aiXander d573407319 update default configs 2024-04-18 09:12:56 -07:00
aiXander 6f5253c1d4 all seems to work 2024-04-18 08:36:14 -07:00
aiXander f5a1de5ffe move token warmup to before lora init 2024-04-18 07:54:38 -07:00
aiXander 6470658e08 fix txt encoder grads 2024-04-18 07:51:43 -07:00
aiXander 2af80788ca fix txt gradients 2024-04-18 07:50:39 -07:00
aiXander 284b934567 fix typo 2024-04-18 07:30:35 -07:00
aiXander 25c9afe587 small tweaks 2024-04-18 07:28:12 -07:00
aiXander 527f075d7a add debug to lr tracking 2024-04-18 07:19:27 -07:00
aiXander a5b7621732 cleanup optimizers class and add learning rate tracking 2024-04-18 07:15:23 -07:00
mayukhdeb 6a6b5df282 accomodate case when optimizer.optimizer_textual_inversion is None and fix incorrect indent 2024-04-18 05:52:06 -07:00
mayukhdeb b4891984aa fix bug where textual inversion was disabling grads for text encoder lora 2024-04-18 05:51:35 -07:00
xander 9ac8d69bfc add experiments for debugging 2024-04-17 18:48:49 +02:00
xander 223969ed84 add std info to plots, change some defaults 2024-04-17 15:52:46 +02:00
xander 57d5b6b762 merge 2024-04-17 14:37:34 +02:00
xander 2a7c53a0c4 updates 2024-04-17 14:36:53 +02:00
mayukhdeb 4f64882f11 remove validation_prompts_raw overwrting 2024-04-17 05:35:12 -07:00
mayukhdeb 1bd9e875fc switch to the new load fn 2024-04-17 05:34:21 -07:00
mayukhdeb bcaccf01c9 new example command 2024-04-17 05:34:01 -07:00
mayukhdeb 0bd726b6be unified save and load functions 2024-04-17 05:33:50 -07:00
mayukhdeb cde4432588 cleaner checkpoint saving WIP 2024-04-17 04:56:34 -07:00
xander aaf090c2e8 fix save/load bug 2024-04-16 21:57:10 +02:00
xander 5181bbdff6 updates to test ti 2024-04-16 21:41:20 +02:00
mayukhdeb 29d62cd705 offload unet optimizer selection 2024-04-16 07:27:00 -07:00
mayukhdeb 9a985de529 better error message 2024-04-16 07:26:38 -07:00
mayukhdeb fd0c8cde12 offload unet lora param stuff 2024-04-16 06:08:12 -07:00
mayukhdeb fcc569c0a6 offload textual inversion stuff 2024-04-16 05:26:35 -07:00
mayukhdeb 7a70c18c05 offload text encoder lora peft stuff into functions 2024-04-16 01:15:36 -07:00
mayukhdeb a4a4d9aaa3 fix indentation bug on text encoder lora params 2024-04-16 00:54:04 -07:00
xander c4ea271482 fix textual_inversion bug 2024-04-16 03:22:10 +02:00
xander 66b34b312e merge 2024-04-16 01:12:24 +02:00
xander 0f9f5a7009 tiny changes 2024-04-16 01:11:28 +02:00
aiXander b8c668822b setup for gridsearch 2024-04-15 16:10:36 -07:00
aiXander 85b86f3667 remove trainable embeddings from print 2024-04-15 08:36:18 -07:00
aiXander 8d58e312ed update deps 2024-04-15 08:05:52 -07:00
aiXander a4b8149ff3 update README 2024-04-15 02:09:11 -07:00
mayukhdeb 4385a55a23 keep all optimizers in one place 2024-04-15 00:25:26 -07:00
mayukhdeb de0d8a8a41 potential fix on the bug where embedding grads were None when text encoder lora == adamw 2024-04-14 11:57:39 -07:00
aiXander b9a315bb98 try to debug bug issue 2024-04-14 11:20:47 -07:00
aiXander c788afa996 try to debug bug issue 2024-04-14 11:07:49 -07:00
aiXander e606340c56 update defaults 2024-04-14 08:09:33 -07:00
mayukhdeb 08c7394677 add new args 2024-04-11 23:04:08 -07:00
mayukhdeb 54ba9492bc auto download checkpoints 2024-04-11 23:01:16 -07:00
mayukhdeb 13ba5da163 fix small bug 2024-04-11 07:11:14 -07:00
mayukhdeb 430c8c62d6 unet full finetuning + checkpointing + eval (commented out gradient norm logging temporarily) 2024-04-11 06:43:01 -07:00
mayukhdeb 9d63a9c1fd fix adapter_config.json not found error 2024-04-10 05:44:29 -07:00
mayukhdeb ff7b53ee63 load text encoder LoRAs for eval 2024-04-10 05:44:13 -07:00
mayukhdeb 744051b7b3 finetune and save text encodder w/ LoRA 2024-04-10 00:22:05 -07:00
mayukhdeb e19b73f3dd add training_image_alignment + remove upper_triangle in clip_diversity 2024-04-09 23:21:11 -07:00
aiXander 501f9991ab fix resolution bug 2024-04-09 17:53:07 -07:00
aiXander 0e5867c4f8 fix training resolution rounding error 2024-04-09 15:38:14 -07:00
aiXander 8d970affc1 first pass to try and fix ComfyUI loading 2024-04-09 15:34:00 -07:00
mayukhdeb 22737bb36c update todos 2024-04-09 03:58:58 -07:00
mayukhdeb bda436a5b7 ignore eval generated images 2024-04-09 03:57:23 -07:00
mayukhdeb 493310e51f move sval scripts to correct folder 2024-04-09 03:56:46 -07:00
mayukhdeb 1ede738580 load aesthetic model checkpoint 2024-04-09 03:56:01 -07:00
mayukhdeb eadb18c0d0 fix eval script 2024-04-09 03:51:09 -07:00
mayukhdeb 6aff99be49 ignore aesthetic model (used for eval) 2024-04-09 02:56:36 -07:00
mayukhdeb 221925c6ef add eval notes 2024-04-09 02:55:40 -07:00
aiXander 2eeecae5d1 add aspect ratio training 2024-04-09 01:48:43 -07:00
aiXander 8ec0060ecd update todos 2024-04-09 00:17:01 -07:00
aiXander 72eae9e1eb update default warmup steps 2024-04-08 23:47:44 -07:00
aiXander a474550c1b remove ti from dataloader, simplify 2024-04-08 23:43:33 -07:00
aiXander 78011776f7 push updates 2024-04-08 23:01:54 -07:00
aiXander c01a99c17f add token warmup 2024-04-08 19:36:54 -07:00
aiXander ceb91f16e2 fix reproducibility seed bug 2024-04-08 02:07:52 -07:00
aiXander f31643c98a tweaks 2024-04-08 01:45:20 -07:00
aiXander 303eb29ece fix tiny bug in predict 2024-04-08 01:30:16 -07:00
aiXander f33743c29d fix cog predict 2024-04-08 01:24:07 -07:00
aiXander 2fe846eba2 final tweaks 2024-04-08 01:07:16 -07:00
aiXander 85748c11f0 final tweaks 2024-04-08 01:04:40 -07:00
aiXander cec51e15d8 big refactor final 2024-04-08 00:59:00 -07:00
aiXander ea470f5de7 cleanup #3 2024-04-08 00:02:08 -07:00
aiXander 868e0e1b2f cleanup #2 2024-04-07 23:36:37 -07:00
aiXander 8bcced05f1 cleanup #1 2024-04-07 23:17:30 -07:00
aiXander 12588e3124 add comment 2024-04-07 19:21:33 -07:00
aiXander 24482e5009 tweaks 2024-04-07 19:19:44 -07:00
aiXander d642537a33 tweaks 2024-04-07 18:54:20 -07:00
aiXander 5fdd0c1172 loss calc cleanup 2024-04-07 18:51:15 -07:00
aiXander 129ba6855c cleanup 2024-04-07 18:13:28 -07:00
aiXander 46390a0802 update README 2024-04-07 17:43:50 -07:00
aiXander cba5273546 more cleanup 2024-04-07 17:21:06 -07:00
aiXander 7970aab689 fix two small bugs 2024-04-07 15:13:12 -07:00
aiXander d9a4e8f361 inference updates 2024-04-06 03:14:52 -07:00
aiXander ad952988db inference updates 2024-04-06 00:28:12 -07:00
aiXander 3870c8ebda release clipseg model after use 2024-04-05 13:56:12 -07:00
aiXander 820edaf20b update default config 2024-04-05 03:33:12 -07:00
aiXander 8230058e65 Merge branch 'bughunt' of https://github.com/edenartlab/trainer into bughunt 2024-04-05 01:46:00 -07:00
aiXander 71fd001f78 make token_amount an input config param 2024-04-05 01:45:54 -07:00
mayukhdeb 51cd996b3c update todos 2024-04-05 01:10:02 -07:00
mayukhdeb 4e9d5073a7 handle face and object concepts 2024-04-05 01:09:42 -07:00
mayukhdeb 0c7c225909 use cossim matrix for diversity thing 2024-04-04 23:24:19 -07:00
aiXander 129e5f756a update regs 2024-04-04 19:33:36 -07:00
aiXander 3a6aa8cec3 massive cleanup 2024-04-04 19:27:33 -07:00
mayukhdeb 13d57aaaed implement image text cossim for style prompts 2024-04-04 02:31:43 -07:00
mayukhdeb bf91ee3e83 return filenames and prompts 2024-04-04 02:31:43 -07:00
aiXander 036a899400 update gridsearch generation 2024-04-04 00:53:08 -07:00
aiXander 386e56a83a update gridsearch generation 2024-04-04 00:52:48 -07:00
aiXander e517bd335b Merge branch 'bughunt' of https://github.com/edenartlab/trainer into bughunt 2024-04-04 00:04:26 -07:00
aiXander eef877e206 push updates and hyperparam script 2024-04-04 00:04:20 -07:00
mayukhdeb f593fdf5a7 rename file 2024-04-03 23:54:09 -07:00
mayukhdeb 7d46a193dc enable aspect ratio bucketing 2024-04-03 23:53:37 -07:00
mayukhdeb 14bfba282a return filenames of images 2024-04-03 23:53:37 -07:00
mayukhdeb 3a55b517f2 handle json stuff 2024-04-03 23:53:37 -07:00
mayukhdeb e43cb3891f add new dependencies for eval 2024-04-03 23:53:37 -07:00
mayukhdeb da9781fd69 eval overhaul wip (missing aesthetic checkpoint) 2024-04-03 23:53:37 -07:00
aiXander a840283f64 more improvements on token embeddings 2024-04-03 15:56:30 -07:00
aiXander 9198b2adab make sd15 work 2024-04-03 03:31:11 -07:00
aiXander 76c445f073 push final changes 2024-04-02 23:12:45 -07:00
aiXander 8d1d4560c8 update std init 2024-04-02 22:14:05 -07:00
aiXander 9dc392a7b9 first pass at fixing sd15 bug 2024-04-02 18:53:13 -07:00
aiXander 660aa3fa16 bugfixes and updates 2024-04-02 18:29:56 -07:00
aiXander 5fb1e6c39c update token embedding plots 2024-04-02 17:23:38 -07:00
mayukhdeb 161d3ad17f wip render_images_eval 2024-04-02 11:36:02 -07:00
mayukhdeb 0161bc5935 less ugly batch loading for aspect ratio bucketing 2024-04-02 06:00:45 -07:00
mayukhdeb c73f6462d2 fix typo 2024-04-02 05:44:29 -07:00
aiXander e277232e62 tweaks 2024-04-02 03:01:00 -07:00
aiXander d404d59952 adjust default training res 2024-04-02 02:19:01 -07:00
aiXander ba28aedd15 push final changes 2024-04-02 02:15:02 -07:00
aiXander f25f40a470 bugfixes, all seems working now 2024-04-02 01:33:39 -07:00
aiXander d28214c850 stash changes before backtracking 2024-04-01 20:33:32 -07:00
aiXander cd62af4049 tweaks to validation img rendering 2024-04-01 15:50:26 -07:00
mayukhdeb 8b20f2da23 imeplement aspect ratio bucketing 2024-04-01 07:33:57 -07:00
aiXander 059f7f2645 fix moved file 2024-04-01 01:05:26 -07:00
aiXander 6007d09366 Merge branch 'bughunt' of https://github.com/edenartlab/trainer into bughunt 2024-04-01 00:56:07 -07:00
aiXander cff6fd00a4 simplify preprocessing 2024-04-01 00:56:00 -07:00
mayukhdeb abead52a2f fix import 2024-03-31 04:47:26 -07:00
mayukhdeb eaaab69d8b move seed fn 2024-03-31 04:47:16 -07:00
mayukhdeb 671ff7b8dc comment out grad norm stuff (it was throwing errors) 2024-03-31 04:47:00 -07:00
mayukhdeb 12ec54e069 remove unused stuff 2024-03-31 04:46:21 -07:00
mayukhdeb cd355d6e7e fix error: cannot import name 'convert_all_state_dict_to_peft' from 'diffusers.utils' 2024-03-31 04:46:07 -07:00
mayukhdeb ff0c34e9da move script 2024-03-31 01:02:33 -07:00
mayukhdeb 65d759dfeb remove unused args 2024-03-31 00:34:04 -07:00
aiXander 67445d3963 small tweaks 2024-03-30 21:17:09 -07:00
aiXander 4e97b176a1 replace TOKEN with TOK 2024-03-30 21:00:03 -07:00
aiXander 1728715cc2 update defaults 2024-03-30 20:56:20 -07:00
aiXander 90feb0709c add auto-gpu picking 2024-03-30 20:37:01 -07:00
aiXander d56af4c5ed minor updates 2024-03-30 20:12:46 -07:00
aiXander e3aa07ccfb add token modulation for inference 2024-03-30 19:55:19 -07:00
aiXander 3f90825f7d big refactor and add inference code 2024-03-30 19:18:16 -07:00
aiXander 6091dd7a2c update saving hooks and optimizer code 2024-03-30 15:41:06 -07:00
aiXander 11fd00e5b0 small tweaks 2024-03-30 14:17:58 -07:00
aiXander 4a3112c5be update default models 2024-03-30 04:59:17 -07:00
aiXander 746279f20b simplify render 2024-03-30 02:36:50 -07:00
aiXander e50844e79f simplify dataloader 2024-03-30 00:14:22 -07:00
aiXander e0de792a9d sample latents from vae distribution at each training step 2024-03-29 23:50:58 -07:00
aiXander 8ce5088064 reduce default grad clip 2024-03-29 23:21:58 -07:00
aiXander be368b78bb add gradient tracking 2024-03-29 22:18:29 -07:00
aiXander b1ab40b15c more config changes and cleanup 2024-03-29 21:24:10 -07:00
aiXander da7dbd6000 update defaults 2024-03-29 18:20:34 -07:00
aiXander b4d6a0dda8 small changes to lora 2024-03-29 16:45:03 -07:00
aiXander cd9b6629bc update args 2024-03-29 15:09:10 -07:00
aiXander 6b8715663d small refactoring, adjusting defaults 2024-03-29 14:35:12 -07:00
mayukhdeb e19713a4b8 add training config file 2024-03-29 12:36:42 -07:00
mayukhdeb 7d8b7b2c63 pass trigger_text as arg 2024-03-29 12:36:28 -07:00
mayukhdeb d2b30450e4 update config args 2024-03-29 12:36:06 -07:00
mayukhdeb 975e385503 add preprocess step + mask grads 2024-03-29 12:35:44 -07:00
mayukhdeb 0a25118461 update todos 2024-03-29 12:34:44 -07:00
mayukhdeb 9d383569bc fix import 2024-03-29 12:34:35 -07:00
mayukhdeb c2468dd46d small fix 2024-03-29 00:52:47 -07:00
mayukhdeb 3ac41fbe14 small cleanup+formatting 2024-03-28 05:33:38 -07:00
mayukhdeb e7e3eecb48 verbose imports 2024-03-28 03:55:54 -07:00
mayukhdeb 61944f27e8 cleanup imports 2024-03-28 03:48:29 -07:00
mayukhdeb 364dfc8157 modularization 2024-03-27 07:11:38 -07:00
mayukhdeb 99187c1f67 update todos 2024-03-27 06:45:56 -07:00
mayukhdeb 904a5e6237 add requirements 2024-03-26 10:48:39 -07:00
mayukhdeb 28c13c21f8 enable training without cog 2024-03-26 10:48:16 -07:00
mayukhdeb 67574ef1b8 config: classmethod to init from filename 2024-03-26 05:33:38 -07:00
mayukhdeb 68e61b6c1d migration to config is complete (working) 2024-03-26 05:27:05 -07:00
mayukhdeb 9d4bc0ae2c more progress (working) 2024-03-26 01:10:32 -07:00
mayukhdeb fac2811fc6 more progress (working) 2024-03-26 00:22:55 -07:00
mayukhdeb fb57183706 big changes (working) 2024-03-25 23:45:56 -07:00
mayukhdeb 63def6ade0 more progress (working) 2024-03-25 23:17:56 -07:00
mayukhdeb 5e893857bb more progress (working) 2024-03-25 07:11:22 -07:00
mayukhdeb 2fef34a4b9 more progress (working) 2024-03-25 05:11:45 -07:00
mayukhdeb 3e5a6cbe35 more progress (things are working) 2024-03-25 04:50:08 -07:00
mayukhdeb 28eff44081 more progress (thinsg working) 2024-03-25 04:26:20 -07:00
mayukhdeb f2764cd309 remove unused args 2024-03-24 23:59:43 -07:00
mayukhdeb ccf010cdb3 continue migration to config (things are working) 2024-03-24 22:51:10 -07:00
mayukhdeb d72005a9c0 slow transition to useing config (everything seems to be working) 2024-03-17 05:40:40 -07:00
mayukhdeb 8547000fdb remove part with unknown variable: model_input 2024-03-17 05:17:12 -07:00
mayukhdeb 21a4c42da0 dtype selection ofload 2024-03-17 05:14:06 -07:00
mayukhdeb 500e974f5e remove unused arg (things are working) 2024-03-17 00:49:01 -07:00
mayukhdeb 4673ec5ef4 nuke args_dict (everything working) 2024-03-16 06:47:55 -07:00
mayukhdeb 0c8c0edbfd get rid of args_dict and initiate slow transition to using config (everything is still working) 2024-03-16 06:20:14 -07:00
mayukhdeb 59538e6995 more args in config (things are still working) 2024-03-16 05:05:46 -07:00
mayukhdeb dc6023b2f8 add method to save as json 2024-03-16 04:39:40 -07:00
mayukhdeb 130db42b77 refactor to work on actual args and not config 2024-03-16 04:38:53 -07:00
mayukhdeb c185b0f0d6 offload stuff from predict.py 2024-03-16 04:38:37 -07:00
mayukhdeb 8326386934 predict: move config to just below main fn + change function names + offload stuff into obtain_inserting_list_tokens 2024-03-16 04:38:19 -07:00
mayukhdeb 039c93ad70 predict: remove positional args 2024-03-15 09:55:52 -07:00
mayukhdeb c95978088c predict: offload download and model info 2024-03-15 09:21:11 -07:00
mayukhdeb 42c31241af modify config based on concept mode 2024-03-15 07:16:43 -07:00
mayukhdeb 09aa083e1b remove old script 2024-03-15 07:09:39 -07:00
mayukhdeb d4916bafc9 predict: move seed fn 2024-03-15 07:08:35 -07:00
mayukhdeb 435d043873 stop ignoring trainer module 2024-03-15 06:40:13 -07:00
mayukhdeb 269eeec1bc slow port to a standard config 2024-03-15 06:27:18 -07:00
54 changed files with 7181 additions and 2772 deletions
+32
View File
@@ -0,0 +1,32 @@
# The .dockerignore file excludes files from the container build process.
# https://docs.docker.com/engine/reference/builder/#dockerignore-file
# Exclude Git files
.git
.github
.gitignore
# Exclude Python cache files
__pycache__
.mypy_cache
.pytest_cache
.ruff_cache
# exclude trained model rars:
*.rar
# Dev folders:
debug
xander
datasets
rendered_images
# trained models:
lora_models/*
# Ignore the entire models folder by default:
models/*
### Include pipeline models: ###
!models/juggernaut_reborn.safetensors
!models/juggernaut_v6.safetensors
+17 -6
View File
@@ -1,13 +1,24 @@
models
lora_models
remove
cache
__pycache__
.ipynb_checkpoints/
models
lora_models*
eden_lora_training_runs/
datasets
*.tar
.env
.cog
xander*.sh
.huggingface
tests/
rendered_images*
gridsearch*
aesthetic_score_best_model.pth
# experiment folders:
conditioning_spaces/
training_args_x_*.json
xander_configs/
debug/*
!debug/*.py
+75 -32
View File
@@ -1,9 +1,42 @@
# Trainer
Code for finetuning and training LoRa modules on top of Stable Diffusion.
This trainer was developed by the [**Eden** team](https://eden.art/)
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
<p align="center">
<strong>Training images:</strong><br>
<img src="assets/xander_training_images.jpg" alt="Image 1" style="width:80%;"/>
</p>
<p align="center">
<strong>Generated imgs with trained LoRa:</strong><br>
<img src="assets/xander_generated_images.jpg" alt="Image 2" style="width:80%;"/>
</p>
The trainer supports 3 default modes:
- **style**: used for learning the aesthetic style of a collection of images.
- **face**: used for learning a specific face (can be human, character, ...).
- **object**: will learn a specific object or thing featured in the training images.
## Setup
Install all dependencies using
`pip install -r requirements.txt`
then you can simply run:
`python main.py train_configs/training_args.json`
to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job.
---
You can also run this through Replicate using cog (~docker image):
1. Install Replicate 'cog':
```
@@ -11,52 +44,62 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
sudo chmod +x /usr/local/bin/cog
```
2. Build the image with `sudo cog build`
3. Run a training run with `sudo sh test_train.sh`
2. Build the image with `cog build`
3. Run a training run with `sh cog_test_train.sh`
4. You can also go into the container with `cog run /bin/bash`
## Automatic Checkpoint Evaluation
This script uses CLIP img/txt similarity scores to evaluate how good the LoRa is vs how overfit.
Download the aesthetic predictor model checkpoint first from google drive. This should give you a file named: `aesthetic_score_best_model.pth` (99.2 MB)
```bash
gdown 1thEIlXVc8lkULVUBY9Ab45tsOERxkjxns
```
Once the model is downloaded, you can run the eval script with the following CLI args:
- `output_folder`: this is where the outputs of the model get saved as jpeg files
- `lora_path`: path to your LoRA checkpoint (make sure you edit `path_to_your_model_checkpoints` to point to the correct folder. It generally ends with something like `checkpoint-600` where `600` was the training step)
- `output_json`: save all scores in this json file
- `config_filename`: config file used for training
```bash
python3 evaluate.py \
--output_folder eval_images \
--lora_path path_to_your_model_checkpoint \
--output_json eval_results.json \
--config_filename training_args.json
```
## TODO's
Code / Cleanup:
- turn all/most of the args of the main() function in trainer_pti.py and the preprocess() function into a clean args_dict that makes it easy to add and distribute new parameters over the code and save these args to a .json file at the end.
- Modularize the logic in train.py as much as possible, trying to minimize dev work that needs to happen when SD3 drops
- make a clean train.py entrypoint that can be run as a normal python command (instead of having to use cog)
- make it so the textual_inversion optimizer only optimizes the actual trained token embeddings instead of all of them + resetting later
- test if the trained concepts with peft are compatible with ComfyUI / AUTO1111
Bugs:
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
Algo:
- Add aspect_ratio bucketing into the dataloader so we can train on non-square images (take this from https://github.com/kohya-ss/sd-scripts)
- Improve some of the chatgpt functionality:
- separate the "gpt_description" / "gpt_segmentation" prompt calls and make them run on a subset of prompts in case there's a lot of imgs / prompts (possibly use img_grids for some gpt4-v calls)
- currently some sub-optimal stuff can happen in preprocess() when there's less than 3 or more than 45 imgs, try to improve this
- Test if timesteps = torch.randint() can be improved: look at sdxl training code! (see https://github.com/huggingface/diffusers/blob/main/examples/advanced_diffusion_training/train_dreambooth_lora_sdxl_advanced.py#L1263, https://arxiv.org/pdf/2206.00364.pdf)
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
- test if textual inversion training can also happen with prodigy_optimizer
- the random initialization of the token embeddings has a relatively large impact on the final outcome, there are prob ways to reduce
this random variance, eg CLIP_similarity pretraining.
- Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM
- it looks the like gpu-utilization is only like 65-70% during training: whats the bottleneck? Can we speed this up?
Bugfixing:
see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037
- check how the pipe() objects work under the hood in HF diffusers library, is there a difference w how the unet is called in the training loop / which args it gets?
- Try to find out why the diffusers training script works for sd15 and ours doesnt:
See here: https://huggingface.co/blog/sdxl_lora_advanced_script
and here: https://github.com/huggingface/diffusers/tree/main/examples/advanced_diffusion_training
- figure out how to adaptively set lora_scale at inference time using peft + diffusers? (https://github.com/huggingface/peft/blob/main/src/peft/tuners/lora/layer.py#L240)
- improve data augmentation, eg by adding outpainted, smaller versions of faces / objects
Small, minor tweaks:
- preprocess.py: the imgs are first auto-captioned and then cropped, this is not ideal, swap this around!
Bigger improvements:
- add stronger token regularization (eg CelebBasis spanning basis):
- Add multi-token training
- implement perfusion: https://research.nvidia.com/labs/par/Perfusion/
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- implement prompt-aligned: https://prompt-aligned.github.io/
- make compatible with ziplora: https://ziplora.github.io/
Tuning Experiments once code is fully ready:
- test if VAE weight_type actually matters for training
Tuning Experiments:
- try-out conditioning noise injection during training to increase robustness
- re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam)
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
but it then recovers. Can we avoid this collapse? Is the learning rate too high?
- gradient_accumulation
- offset noise
- AB test Dora vs Lora
- sweep n_trainable_tokens to inject
+10
View File
@@ -0,0 +1,10 @@
import os
import sys
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from node import Eden_LoRa_trainer
NODE_CLASS_MAPPINGS = {
"Eden_LoRa_trainer": Eden_LoRa_trainer,
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 915 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 342 KiB

+5 -27
View File
@@ -3,36 +3,14 @@
build:
gpu: true
cuda: "11.8"
python_version: "3.9"
cuda: "12.1"
python_version: "3.11"
system_packages:
- "libgl1-mesa-glx"
- "ffmpeg"
- "libsm6"
- "libxext6"
python_packages:
- "scipy"
- "diffusers==0.25.1"
- "peft==0.9.0"
- "torch==2.0.1"
- "transformers==4.31.0"
- "invisible-watermark==0.2.0"
- "accelerate==0.21.0"
- "pandas==2.0.3"
- "torchvision==0.15.2"
- "numpy==1.25.1"
- "pandas==2.0.3"
- "fire==0.5.0"
- "opencv-python>=4.1.0.25"
- "mediapipe==0.10.2"
- "openai==1.2.4"
- python-dotenv
- prodigyopt
- omegaconf
python_requirements: requirements.txt
run:
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/download/v0.0.1/pget" && chmod +x /usr/local/bin/pget
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
predict: "predict.py:Predictor"
image: "r8.im/abraham-ai/sdxl-lora-trainer"
image: "r8.im/edenartlab/sdxl-lora-trainer"
+12
View File
@@ -0,0 +1,12 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=3"
cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip" \
-i concept_mode="face" \
-i sd_model_version="sdxl" \
-i max_train_steps="360" \
-i caption_model="blip" \
-i debug="False" \
-i seed="0"
File diff suppressed because it is too large Load Diff
+533
View File
@@ -0,0 +1,533 @@
{
"last_node_id": 12,
"last_link_id": 23,
"nodes": [
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 16
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
6
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"text, watermark"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1209,
188
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 8
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
9
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 12,
"type": "Reroute",
"pos": [
220,
14
],
"size": [
75,
26
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 22
}
],
"outputs": [
{
"name": "",
"type": "MODEL",
"links": [
19
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 3,
"type": "KSampler",
"pos": [
863,
186
],
"size": {
"0": 315,
"1": 262
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 19
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 4
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 6
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
7
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
1,
"fixed",
25,
8,
"euler",
"normal",
1
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
473,
609
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
768,
768,
1
]
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-467,
120
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
10
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
12
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
8
],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"juggernaut_reborn.safetensors"
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
415,
186
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 15
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
4
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"a photo of embedding:xander_sd15_embedding on the beach "
]
},
{
"id": 11,
"type": "Reroute",
"pos": [
220,
49
],
"size": [
75,
26
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 23
}
],
"outputs": [
{
"name": "",
"type": "CLIP",
"links": [
15,
16
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 9,
"type": "SaveImage",
"pos": [
347,
-270
],
"size": {
"0": 210,
"1": 270
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 9
}
],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 10,
"type": "LoraLoader",
"pos": [
-72,
-229
],
"size": {
"0": 254.95774841308594,
"1": 127.86701202392578
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 10
},
{
"name": "clip",
"type": "CLIP",
"link": 12
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
22
],
"shape": 3,
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
23
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LoraLoader"
},
"widgets_values": [
"xander_sd15_lora.safetensors",
0.6,
0.6
]
}
],
"links": [
[
2,
5,
0,
3,
3,
"LATENT"
],
[
4,
6,
0,
3,
1,
"CONDITIONING"
],
[
6,
7,
0,
3,
2,
"CONDITIONING"
],
[
7,
3,
0,
8,
0,
"LATENT"
],
[
8,
4,
2,
8,
1,
"VAE"
],
[
9,
8,
0,
9,
0,
"IMAGE"
],
[
10,
4,
0,
10,
0,
"MODEL"
],
[
12,
4,
1,
10,
1,
"CLIP"
],
[
15,
11,
0,
6,
0,
"CLIP"
],
[
16,
11,
0,
7,
0,
"CLIP"
],
[
19,
12,
0,
3,
0,
"MODEL"
],
[
22,
10,
0,
12,
0,
"*"
],
[
23,
10,
1,
11,
0,
"*"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": {
"0": 513.8734070325743,
"1": 351.4824273966635
}
}
},
"version": 0.4
}
Executable
+656
View File
@@ -0,0 +1,656 @@
import fnmatch
import math
import os
import time
import shutil
import gc
import numpy as np
import argparse
import itertools
import zipfile
import torch
import torch.utils.checkpoint
from tqdm import tqdm
from typing import Union, Iterable, List, Dict, Tuple, Optional, cast
#from diffusers.training_utils import cast_training_params
from trainer.utils.utils import *
from trainer.checkpoint import save_checkpoint
from trainer.embedding_handler import TokenEmbeddingsHandler
from trainer.dataset import PreprocessedDataset
from trainer.config import TrainingConfig
from trainer.models import print_trainable_parameters, load_models
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer
from trainer.inference import render_images, get_conditioning_signals
from trainer.preprocess import preprocess
from trainer.utils.io import make_validation_img_grid
from trainer.optimizer import (
OptimizerCollection,
get_optimizer_and_peft_models_text_encoder_lora,
get_textual_inversion_optimizer,
get_unet_lora_parameters,
get_unet_optimizer
)
def train(config: TrainingConfig):
seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
from trainer.ti_cross_attn_loss import init_daam_loss
pipe, daam_loss = init_daam_loss(
pipeline=pipe
)
config.sd_model_version = sd_model_version
config.pretrained_model["version"] = sd_model_version
config, input_dir = preprocess(
config,
working_directory=config.output_dir,
concept_mode=config.concept_mode,
input_zip_path=config.lora_training_urls,
caption_text=config.caption_prefix,
mask_target_prompts=config.mask_target_prompts,
target_size=config.resolution,
crop_based_on_salience=config.crop_based_on_salience,
use_face_detection_instead=config.use_face_detection_instead,
left_right_flip_augmentation=config.left_right_flip_augmentation,
augment_imgs_up_to_n = config.augment_imgs_up_to_n,
caption_model = config.caption_model,
seed = config.seed,
)
if config.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
text_encoders = [text_encoder_one, text_encoder_two],
tokenizers = [tokenizer_one, tokenizer_two]
)
embedding_handler.initialize_new_tokens(
inserting_toks=config.inserting_list_tokens,
starting_toks=None,
seed=config.seed
)
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
embedding_handler.make_embeddings_trainable()
embedding_handler.token_regularizer = ConditioningRegularizer(config, embedding_handler)
embedding_handler.pre_optimize_token_embeddings(config, pipe)
# Turn off all gradients for now:
unet.requires_grad_(False)
vae.requires_grad_(False)
text_encoders = embedding_handler.text_encoders
for txt_encoder in text_encoders:
if txt_encoder is not None:
txt_encoder.requires_grad_(False)
if config.text_encoder_lora_optimizer is not None:
print("Creating LoRA for text encoder...")
optimizer_text_encoder_lora , text_encoder_peft_models = get_optimizer_and_peft_models_text_encoder_lora(
text_encoders=text_encoders,
lora_rank = config.text_encoder_lora_rank,
lora_alpha_multiplier = config.lora_alpha_multiplier,
use_dora = config.use_dora,
optimizer_name = config.text_encoder_lora_optimizer,
lora_lr = config.text_encoder_lora_lr,
weight_decay = config.text_encoder_lora_weight_decay
)
else:
optimizer_text_encoder_lora = None
text_encoder_peft_models = [None] * len(text_encoders)
embedding_handler.make_embeddings_trainable()
if not config.disable_ti:
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
else:
optimizer_ti = None
textual_inversion_params = None
if not config.is_lora: # This code pathway has not been tested in a long while
print(f"Doing full fine-tuning on the U-Net")
unet.requires_grad_(True)
unet_lora_parameters = None
optimizer_text_encoder_lora = None
unet_trainable_params = unet.parameters()
else:
# Do lora-training instead.
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
# target_blocks=["block"] for original IP-Adapter
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
unet, unet_trainable_params, unet_lora_parameters = get_unet_lora_parameters(
lora_rank = config.lora_rank,
lora_alpha_multiplier = config.lora_alpha_multiplier,
lora_weight_decay=config.lora_weight_decay,
use_dora = config.use_dora,
unet=unet,
pipe=pipe
)
optimizer_unet = get_unet_optimizer(
prodigy_d_coef=config.prodigy_d_coef,
prodigy_growth_factor=config.unet_prodigy_growth_factor,
lora_weight_decay=config.lora_weight_decay,
use_dora=config.use_dora,
unet_trainable_params=unet_trainable_params,
optimizer_name=config.unet_optimizer_type
)
print_trainable_parameters(unet, model_name = 'unet')
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
print_trainable_parameters(text_encoder, model_name = f'text_encoder_{i}')
train_dataset = PreprocessedDataset(
input_dir,
pipe,
vae.float(),
size = config.train_img_size,
do_cache=config.do_cache,
substitute_caption_map=config.token_dict,
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
train_batch_size=config.train_batch_size
)
# offload the vae to cpu:
vae = vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=config.train_batch_size,
shuffle=True,
num_workers=config.dataloader_num_workers,
)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
if config.max_train_steps is None:
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
print(f"--- Num samples = {len(train_dataset)}")
print(f"--- Num batches each epoch = {len(train_dataloader)}")
print(f"--- Num Epochs = {config.num_train_epochs}")
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
global_step = 0
last_save_step = 0
progress_bar = tqdm(range(global_step, config.max_train_steps), position=0, leave=True)
checkpoint_dir = os.path.join(str(config.output_dir), "checkpoints")
if os.path.exists(checkpoint_dir):
shutil.rmtree(checkpoint_dir)
os.makedirs(f"{checkpoint_dir}")
# Data tracking inits:
start_time, images_done = time.time(), 0
prompt_embeds_norms = {'main':[], 'reg':[]}
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': []}
grad_norms, token_stds = {'unet': []}, {}
for i in range(len(text_encoders)):
grad_norms[f'text_encoder_{i}'] = []
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
# default value of cold (pre-warmup) optimizer lr:
if config.sd_model_version == "sdxl":
if config.is_lora: # let textual_inversion do the work first!
base_lr = 1.0e-5
else:
base_lr = 3.0e-5
elif config.sd_model_version == "sd15":
# let lora training kick in soonish (pure ti for sd15 is not working super well in my tests)
base_lr = 1.0e-4
#######################################################################################################
"""
Storing all optimizers in a single container
"""
optimizer_collection = OptimizerCollection(
optimizer_textual_inversion=optimizer_ti,
optimizer_text_encoders=optimizer_text_encoder_lora,
optimizer_unet=optimizer_unet,
debug = config.debug
)
optimizers = optimizer_collection.optimizers
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
for epoch in range(config.num_train_epochs):
if config.aspect_ratio_bucketing:
train_dataset.bucket_manager.start_epoch()
progress_bar.set_description(f"# Trainer step: {global_step}, epoch: {epoch}")
for step, batch in enumerate(train_dataloader):
progress_bar.update(1)
finegrained_epoch = epoch + step / len(train_dataloader)
completion_f = finegrained_epoch / config.num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training
if config.ti_optimizer != "prodigy": # Update ti_learning rate gradually:
if optimizers['textual_inversion'] is not None:
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0
# warmup the ti-lr:
if config.ti_lr_warmup_steps > 0:
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0)
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
if config.freeze_ti_after_completion_f <= completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
if optimizers['text_encoders'] is not None:
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
# warmup the txt-encoder lr:
if config.txt_encoders_lr_warmup_steps > 0 and optimizers['text_encoders'] is not None:
warmup_f = min(global_step / config.txt_encoders_lr_warmup_steps, 1.0)
optimizers['text_encoders'].param_groups[0]['lr'] *= warmup_f
if optimizers['unet'] is not None:
# Calculate the exponential factor
exp_factor = (config.unet_lr / base_lr) ** (global_step / config.unet_lr_warmup_steps)
# Apply the exponential learning rate
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
if not config.aspect_ratio_bucketing:
captions, vae_latent, mask = batch
else:
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
captions = list(captions)
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
config, pipe, captions
)
# Sample noise that we'll add to the latents:
vae_latent = vae_latent.to(weight_dtype)
noise = torch.randn_like(vae_latent)
if config.noise_offset > 0.0:
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
noise += config.noise_offset * torch.randn(
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
timesteps = torch.randint(
0,
noise_scheduler.config.num_train_timesteps,
(vae_latent.shape[0],),
device=vae_latent.device,
).long()
noisy_latent = noise_scheduler.add_noise(vae_latent, noise, timesteps)
# Predict the noise residual
model_pred = unet(
noisy_latent,
timesteps,
encoder_hidden_states=prompt_embeds,
timestep_cond=None,
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
return_dict=False,
)[0]
"""
distirbution shift loss
"""
non_ti_heatmaps = []
ti_heatmaps = []
ti_token_indices = [0,1]
batch_index = 0
token_strings = [
pipe.tokenizer.decode(x)
for x in pipe.tokenizer.encode(captions[batch_index])
]
for text_token_index in range(1, len(token_strings)-1):
if text_token_index in ti_token_indices:
ti_heatmaps.append(
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
)
else:
# we unsqueeze because we'll stack them together and then calculate the min, max and the mean
non_ti_heatmaps.append(
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
)
non_ti_heatmaps = torch.cat(
non_ti_heatmaps,
dim = 0
)
ti_heatmaps = torch.cat(
ti_heatmaps,
dim = 0
)
non_ti_dist = {
"mean": non_ti_heatmaps.mean(),
"min": non_ti_heatmaps.min(),
"max": non_ti_heatmaps.min()
}
ti_dist = {
"mean": ti_heatmaps.mean(),
"min": ti_heatmaps.min(),
"max": ti_heatmaps.min()
}
dist_loss = (non_ti_dist["mean"] - ti_dist["mean"].to(non_ti_dist["mean"].device)) ** 2
if global_step % 20 == 0:
batch_index = 0
folder = "./heatmaps"
fig = plt.figure()
token_strings = [
pipe.tokenizer.decode(x)
for x in pipe.tokenizer.encode(captions[batch_index])
]
plot_token_indices = range(len(token_strings))
fig, ax = plt.subplots(nrows=1, ncols=len(plot_token_indices), figsize = (int(3 * len(plot_token_indices)) , 10))
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
im = ax[idx].imshow(heatmap)
ax[idx].set_title(f"{token_strings[text_token_index]}\n timestep: {timesteps[batch_index].item()}\nmax: {heatmap.max().item()}\nmin: {heatmap.min().item()}\nnorm: {heatmap.norm().item()}")
ax[idx].axis("off")
fig.savefig(
os.path.join(
folder,
f"{global_step}.jpg"
)
)
plt.close(fig)
"""
histogram to visualize the distributions of the cross attention values for each text token on the image space
"""
fig = plt.figure()
fig.suptitle(f"Dist loss: {dist_loss.item()}")
plot_token_indices = range(1, len(token_strings)-1)
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
plt.hist(heatmap.reshape(-1), bins = 30, label = token_strings[text_token_index], alpha = 0.5)
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
plt.xlabel("Value")
plt.ylabel("Number of instances")
plt.grid()
# Adjust the layout to prevent the legend from being cut off
plt.tight_layout()
fig.savefig(
os.path.join(
folder,
f"{global_step}_heatmap.jpg"
),
bbox_inches='tight' # This ensures the legend is not cut off when saving
)
plt.close(fig) # Close the figure to free up memory
# Compute the loss:
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
losses['img_loss'].append(loss.item())
if config.training_attributes["gpt_description"] and config.debug:
concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe)
# Dont apply this loss, just plot it for now:
loss += 0.0 * concept_description_loss
losses['concept_description_loss'].append(concept_description_loss.item())
if config.l1_penalty > 0.0 and unet_lora_parameters:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += config.l1_penalty * l1_norm
if optimizers['textual_inversion'] is not None and optimizers['textual_inversion'].param_groups[0]['lr'] > 0.0:
loss, losses, prompt_embeds_norms = embedding_handler.token_regularizer.apply_regularization(loss, losses, prompt_embeds_norms, prompt_embeds, pipe = pipe)
losses['tot_loss'].append(loss.item())
loss = loss + 1e-4 * dist_loss
loss = loss / config.gradient_accumulation_steps
loss.backward()
last_batch = (step + 1 == len(train_dataloader))
if (step + 1) % config.gradient_accumulation_steps == 0 or last_batch:
if optimizers['textual_inversion'] is not None:
# zero out the gradients of the non-trained text-encoder embeddings
for i, embedding_tensor in enumerate(textual_inversion_params):
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
if config.debug:
# Track the average gradient norms:
grad_norms['unet'].append(compute_grad_norm(itertools.chain(unet.parameters())).item())
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
text_encoder_norm = compute_grad_norm(itertools.chain(text_encoder.parameters())).item()
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
optimizer_collection.step()
optimizer_collection.zero_grad()
#############################################################################################################
if config.debug:
# Track the token embedding stds:
trainable_embeddings, _ = embedding_handler.get_trainable_embeddings()
for idx in range(len(text_encoders)):
if text_encoders[idx] is not None:
embedding_stds = trainable_embeddings[f'txt_encoder_{idx}'].detach().float().std(dim=1)
for std_i, std in enumerate(embedding_stds):
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
# Print some statistics:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
config.save_as_json(
os.path.join(output_save_dir, "training_args.json")
)
save_checkpoint(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=config.token_dict,
is_lora=config.is_lora,
unet_lora_parameters=unet_lora_parameters,
name=config.name,
text_encoder_peft_models=text_encoder_peft_models,
pretrained_model_version=config.pretrained_model["version"]
)
last_save_step = global_step
if config.debug:
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings()
for idx, text_encoder in enumerate(text_encoders):
if text_encoder is None:
continue
n = len(token_embeddings[f'txt_encoder_{idx}'])
for i in range(n):
token = trainable_tokens[f'txt_encoder_{idx}'][i]
# Strip any backslashes from the token name:
token = token.replace("/", "_")
embedding = token_embeddings[f'txt_encoder_{idx}'][i]
plot_torch_hist(embedding, global_step, os.path.join(config.output_dir, 'ti_embeddings') , f"enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'red')
embedding_handler.print_token_info()
if config.is_lora: # plotting this hist for full unet parameters can run OOM
plot_torch_hist(unet_lora_parameters, global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
validation_prompts = render_images(
pipe = pipe,
render_size = config.validation_img_size,
lora_path = output_save_dir,
train_step = global_step,
seed = config.seed,
is_lora = config.is_lora,
pretrained_model = config.pretrained_model,
lora_scale = config.sample_imgs_lora_scale,
n_imgs = config.n_sample_imgs,
device = config.device,
checkpoint_folder = None
)
img_grid_path = make_validation_img_grid(output_save_dir)
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
gc.collect()
torch.cuda.empty_cache()
images_done += config.train_batch_size
global_step += 1
if global_step % (config.max_train_steps//50) == 0:
progress = (global_step / config.max_train_steps) + 0.05
#print_system_info()
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
yield np.min((progress, 1.0))
if global_step > config.max_train_steps:
print("Reached max steps, stopping training!", flush = True)
break
# final_save
if (global_step - last_save_step) > 26:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
else:
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
if config.debug:
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
if not os.path.exists(output_save_dir):
os.makedirs(output_save_dir, exist_ok=True)
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
save_checkpoint(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=config.token_dict,
is_lora=config.is_lora,
unet_lora_parameters=unet_lora_parameters,
name=config.name,
pretrained_model_version=config.pretrained_model["version"]
)
if config.debug and 0:
# Reload the entire pipe from disk + LoRa:
pipe_to_use = None
checkpoint_folder = output_save_dir
del unet
del vae
del text_encoder_one
del text_encoder_two
del tokenizer_one
del tokenizer_two
del embedding_handler
del pipe
del train_dataloader
del train_dataset
gc.collect()
torch.cuda.empty_cache()
else:
# Just render images with the active pipe (faster, easier):
pipe_to_use = pipe
checkpoint_folder = None
validation_prompts = render_images(
pipe = pipe_to_use,
render_size=config.validation_img_size,
lora_path=output_save_dir,
train_step=global_step,
seed=config.seed,
is_lora=config.is_lora,
pretrained_model=config.pretrained_model,
lora_scale=config.sample_imgs_lora_scale,
n_imgs = config.n_sample_imgs,
n_steps = 30,
device = config.device,
checkpoint_folder=checkpoint_folder
)
img_grid_path = make_validation_img_grid(output_save_dir)
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
else:
print(f"Skipping final save, {output_save_dir} already exists")
if config.debug:
# Create a zipfile of all the *.py files in the directory
parent_dir = os.path.dirname(os.path.abspath(__file__))
zip_file_path = os.path.join(config.output_dir, 'source_code.zip')
with zipfile.ZipFile(zip_file_path, 'w', zipfile.ZIP_DEFLATED) as zipf:
zipdir(parent_dir, zipf)
config.job_time = time.time() - config.start_time
config.training_attributes["validation_prompts"] = validation_prompts
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
print("Training job complete, saving outputs...", flush = True)
print("------------------------------------------")
return config, output_save_dir
if __name__ == "__main__":
parser = argparse.ArgumentParser(description='Train a concept')
parser.add_argument('config_filename', type=str, help='Input JSON configuration file')
args = parser.parse_args()
config = TrainingConfig.from_json(file_path=args.config_filename)
print("Starting new LoRa training run with config:")
print(config)
print("------------------------------------------")
for progress in train(config=config):
print(f"Progress: {(100*progress):.2f}%", end="\r")
print("Training done :)")
+130
View File
@@ -0,0 +1,130 @@
import os
import tarfile
import json
import time
import torch
import numpy as np
from PIL import Image
from main import train
from trainer.config import TrainingConfig, model_paths
from trainer.utils.io import clean_filename
import folder_paths
import comfy.utils
class Eden_LoRa_trainer:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"training_images_folder_path": ("STRING", {"default": "."}),
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"lora_name": ("STRING", {"default": "Eden_LoRa"}),
"mode": (["style", "face", "object"], ),
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
"use_dora": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}),
"debug_mode": ("BOOLEAN", {"default": False}),
"checkpointing_steps": ("INT", {"default": 200, "min": 10, "max": 2000}),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
}
}
CATEGORY = "Eden 🌱"
RETURN_TYPES = ("IMAGE", "STRING", "STRING", "STRING")
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
FUNCTION = "train_lora"
def train_lora(self,
training_images_folder_path,
ckpt_name,
lora_name = "eden_lora",
mode = "style",
seed = 0,
resolution = 521,
train_batch_size = 4,
max_train_steps = 400,
ti_lr = 0.001,
unet_lr = 0.001,
lora_rank = 16,
use_dora = False,
n_tokens = 2,
debug_mode = False,
checkpointing_steps = 1000,
):
print("Starting new training job...")
# Overwrite hardcoded paths to point to comfyUI folders:
model_paths.set_path("CLIP", os.path.join(folder_paths.models_dir, "clipseg"))
model_paths.set_path("BLIP", os.path.join(folder_paths.models_dir, "blip"))
model_paths.set_path("SR", os.path.join(folder_paths.models_dir, "upscale_models"))
model_paths.set_path("SD", os.path.join(folder_paths.models_dir, "checkpoints"))
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
config = TrainingConfig(
name=lora_name,
lora_training_urls=training_images_folder_path,
concept_mode=mode,
ckpt_path=ckpt_path,
seed=seed,
resolution=resolution,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=checkpointing_steps,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
use_dora=use_dora,
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
debug=debug_mode,
)
pbar = comfy.utils.ProgressBar(100)
with torch.inference_mode(False):
train_generator = train(config=config)
while True:
try:
progress_f = next(train_generator)
pbar.update_absolute(progress_f * 100)
except StopIteration as e:
config, output_save_dir = e.value # Capture the return value
break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
print(f"LORA training node finished in {config.job_time:.1f} seconds")
print("---------- Made with love by Eden.art 🌱 ----------")
# safetensors paths:
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
# find the index of the path containing "_embeddings.safetensors":
for i, path in enumerate(paths):
if "_embeddings.safetensors" in path:
embedding_path = path
else:
lora_path = path
# Load the grid image:
grid_image = Image.open(validation_grid_img_path)
grid_image = np.array(grid_image).astype(np.float32) / 255.0
grid_image = torch.from_numpy(grid_image)[None,]
final_msg = f"LoRa trained in {config.job_time/60:.1f} minutes. Files saved at {output_save_dir}"
return (grid_image, lora_path, embedding_path, final_msg)
+72 -353
View File
@@ -7,16 +7,20 @@ import random
import torch
import numpy as np
import pandas as pd
from collections import OrderedDict
from cog import BasePredictor, BaseModel, File, Input, Path as cogPath
from dotenv import load_dotenv
from preprocess import preprocess
from trainer_pti import main
from main import train
from typing import Iterator, Optional
from io_utils import MODEL_INFO, download_weights, clean_filename
from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig
from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything
DEBUG_MODE = False
XANDER_EXPERIMENT = False
load_dotenv()
@@ -25,6 +29,7 @@ os.environ["TRANSFORMERS_CACHE"] = "/src/.huggingface/"
os.environ["DIFFUSERS_CACHE"] = "/src/.huggingface/"
os.environ["HF_HOME"] = "/src/.huggingface/"
class CogOutput(BaseModel):
files: Optional[list[cogPath]] = []
name: Optional[str] = None
@@ -47,355 +52,93 @@ class Predictor(BasePredictor):
default="unnamed"
),
lora_training_urls: str = Input(
description="Training images for new LORA concept (can be image urls or a .zip file of images)",
default=None
description="Training images for new LORA concept (can be image urls or an url to a .zip file of images)"
),
concept_mode: str = Input(
description=" 'face' / 'style' / 'object' (default)",
default="object",
description="What are you trying to learn?",
choices=["style", "face", "object"],
default="style",
),
sd_model_version: str = Input(
description=" 'sdxl' / 'sd15' ",
description="SDXL gives much better LoRa's if you just need static images. If you want to make AnimateDiff animations, train an SD15 lora.",
choices=["sdxl", "sd15"],
default="sdxl",
),
max_train_steps: int = Input(
description="Number of training steps. Increasing this usually leads to overfitting, only viable if you have > 100 training imgs. For faces you may want to reduce to eg 300",
default=400
),
resolution: int = Input(
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
default=512
),
train_batch_size: int = Input(
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
default=4
),
unet_lr: float = Input(
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
default=0.001
),
ti_lr: float = Input(
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
default=0.001
),
lora_rank: int = Input(
description="Rank of LoRA embeddings for the unet.",
default=16
),
use_dora: bool = Input(
description="Use Dora instead of LoRa",
default=False,
),
n_tokens: int = Input(
description="How many new tokens to train (highly recommended to leave this at 2)",
ge=1, le=3, default=2
),
seed: int = Input(
description="Random seed for reproducible training. Leave empty to use a random seed",
default=None,
),
resolution: int = Input(
description="Square pixel resolution which your images will be resized to for training recommended [768-1024]",
default=960,
),
train_batch_size: int = Input(
description="Batch size (per device) for training",
default=4,
),
num_train_epochs: int = Input(
description="Number of epochs to loop through your training dataset",
default=10000,
),
max_train_steps: int = Input(
description="Number of individual training steps. Takes precedence over num_train_epochs",
default=600,
),
checkpointing_steps: int = Input(
description="Number of steps between saving checkpoints. Set to very very high number to disable checkpointing, because you don't need one.",
default=10000,
),
gradient_accumulation_steps: int = Input(
description="Number of training steps to accumulate before a backward pass. Effective batch size = gradient_accumulation_steps * batch_size",
default=1,
),
is_lora: bool = Input(
description="Whether to use LoRA training. If set to False, will use Full fine tuning",
default=True,
),
prodigy_d_coef: float = Input(
description="Multiplier for internal learning rate of Prodigy optimizer",
default=0.5,
),
ti_lr: float = Input(
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
default=1e-3,
),
ti_weight_decay: float = Input(
description="weight decay for textual inversion embeddings. Don't alter unless you know what you're doing.",
default=3e-4,
),
lora_weight_decay: float = Input(
description="weight decay for lora parameters. Don't alter unless you know what you're doing.",
default=0.002,
),
l1_penalty: float = Input(
description="Sparsity penalty for the LoRA matrices, possibly improves merge-ability and generalization",
default=0.1,
),
lora_param_scaler: float = Input(
description="Multiplier for the starting weights of the lora matrices",
default=0.5,
),
snr_gamma: float = Input(
description="see https://arxiv.org/pdf/2303.09556.pdf, set to None to disable snr training",
default=5.0,
),
lora_rank: int = Input(
description="Rank of LoRA embeddings. For faces 5 is good, for complex concepts / styles you can try 8 or 12",
default=12,
),
caption_prefix: str = Input(
description="Prefix text prepended to automatic captioning. Must contain the 'TOK'. Example is 'a photo of TOK, '. If empty, chatgpt will take care of this automatically",
default="",
),
caption_model: str = Input(
description="Which captioning model to use. ['gpt4-v', 'blip'] are supported right now",
default="blip",
),
left_right_flip_augmentation: bool = Input(
description="Add left-right flipped version of each img to the training data, recommended for most cases. If you are learning a face, you prob want to disable this",
default=True,
),
augment_imgs_up_to_n: int = Input(
description="Apply data augmentation (no lr-flipping) until there are n training samples (0 disables augmentation completely)",
default=20,
),
n_tokens: int = Input(
description="How many new tokens to inject per concept",
default=2,
),
mask_target_prompts: str = Input(
description="Prompt that describes most important part of the image, will be used for CLIP-segmentation. For example, if you are learning a person 'face' would be a good segmentation prompt",
default=None,
),
crop_based_on_salience: bool = Input(
description="If you want to crop the image to `target_size` based on the important parts of the image, set this to True. If you want to crop the image based on face detection, set this to False",
default=True,
),
use_face_detection_instead: bool = Input(
description="If you want to use face detection instead of CLIPSeg for masking. For face applications, we recommend using this option.",
default=False,
),
clipseg_temperature: float = Input(
description="How blurry you want the CLIPSeg mask to be. We recommend this value be something between `0.5` to `1.0`. If you want to have more sharp mask (but thus more errorful), you can decrease this value.",
default=0.7,
),
verbose: bool = Input(description="verbose output", default=True),
run_name: str = Input(
description="Subdirectory where all files will be saved",
default=str(int(time.time())),
),
debug: bool = Input(
description="for debugging locally only (dont activate this on replicate)",
default=False,
),
hard_pivot: bool = Input(
description="Use hard freeze for ti_lr. If set to False, will use soft transition of learning rates",
default=False,
),
off_ratio_power: float = Input(
description="How strongly to correct the embedding std vs the avg-std (0=off, 0.05=weak, 0.1=standard)",
default=0.1,
),
) -> Iterator[GENERATOR_OUTPUT_TYPE]:
"""
lambda @1024 training speed (SDXL):
lambda training speed (SDXL):
bs=2: 3.5 imgs/s, 1.8 batches/s
bs=3: 5.1 imgs/s
bs=4: 6.0 imgs/s,
bs=6: 8.0 imgs/s,
"""
start_time = time.time()
out_root_dir = "lora_models"
if seed is None:
seed = np.random.randint(0, 2**32 - 1)
# Try to make the training reproducible:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
if concept_mode == "face":
left_right_flip_augmentation = False # always disable lr flips for face mode!
mask_target_prompts = "face"
clipseg_temperature = 0.4
if concept_mode == "concept": # gracefully catch any old versions of concept_mode
concept_mode = "object"
if concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
l1_penalty = 0.05
print(f"cog:predict:train_lora:{concept_mode}")
debug = False
print("cog:predict starting new training job...")
if not debug:
yield CogOutput(name=name, progress=0.0)
# Initialize pretrained_model dictionary
pretrained_model = {"version": sd_model_version}
pretrained_model.update(MODEL_INFO[pretrained_model['version']])
# Download the weights if they don't exist locally
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
# hardcoded for now:
token_list = [f"TOK:{n_tokens}"]
#token_list = ["TOK1:2", "TOK2:2"]
token_dict = OrderedDict({})
all_token_lists = []
running_tok_cnt = 0
for token in token_list:
token_name, n_tok = token.split(":")
n_tok = int(n_tok)
special_tokens = [f"<s{i + running_tok_cnt}>" for i in range(n_tok)]
token_dict[token_name] = "".join(special_tokens)
all_token_lists.extend(special_tokens)
running_tok_cnt += n_tok
if 0:
# overwrite some settings for experimentation:
lora_param_scaler = 0.1
l1_penalty = 0.2
prodigy_d_coef = 0.2
ti_lr = 1e-3
lora_rank = 24
lora_training_urls = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/plantoid_5.zip"
concept_mode = "object"
mask_target_prompts = ""
left_right_flip_augmentation = True
output_dir1 = os.path.join(out_root_dir, run_name + "_xander")
input_dir1, n_imgs1, trigger_text1, segmentation_prompt1, captions1 = preprocess(
output_dir1,
concept_mode,
input_zip_path=lora_training_urls,
caption_text=caption_prefix,
mask_target_prompts=mask_target_prompts,
target_size=resolution,
crop_based_on_salience=crop_based_on_salience,
use_face_detection_instead=use_face_detection_instead,
temp=clipseg_temperature,
left_right_flip_augmentation=left_right_flip_augmentation,
augment_imgs_up_to_n = augment_imgs_up_to_n,
seed = seed,
caption_model = caption_model
)
lora_training_urls = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene_5.zip"
concept_mode = "face"
mask_target_prompts = "face"
left_right_flip_augmentation = False
output_dir2 = os.path.join(out_root_dir, run_name + "_gene")
input_dir2, n_imgs2, trigger_text2, segmentation_prompt2, captions2 = preprocess(
output_dir2,
concept_mode,
input_zip_path=lora_training_urls,
caption_text=caption_prefix,
mask_target_prompts=mask_target_prompts,
target_size=resolution,
crop_based_on_salience=crop_based_on_salience,
use_face_detection_instead=use_face_detection_instead,
temp=clipseg_temperature,
left_right_flip_augmentation=left_right_flip_augmentation,
augment_imgs_up_to_n = augment_imgs_up_to_n,
seed = seed,
)
# Merge the two preprocessing steps:
n_imgs = n_imgs1 + n_imgs2
captions = captions1 + captions2
trigger_text = trigger_text1
segmentation_prompt = segmentation_prompt1
# Create merged outdir:
output_dir = os.path.join(out_root_dir, run_name + "_combined")
input_dir = os.path.join(output_dir, "images_out")
os.makedirs(input_dir, exist_ok=True)
# Merge the two preprocessed datasets:
merge_datasets(input_dir1, input_dir2, input_dir, token_dict.keys())
else: # normal, single token run:
output_dir = os.path.join(out_root_dir, run_name)
input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
output_dir,
concept_mode,
input_zip_path=lora_training_urls,
caption_text=caption_prefix,
mask_target_prompts=mask_target_prompts,
target_size=resolution,
crop_based_on_salience=crop_based_on_salience,
use_face_detection_instead=use_face_detection_instead,
temp=clipseg_temperature,
left_right_flip_augmentation=left_right_flip_augmentation,
augment_imgs_up_to_n = augment_imgs_up_to_n,
seed = seed,
)
if not debug:
yield CogOutput(name=name, progress=0.05)
# Make a dict of all the arguments and save it to args.json:
args_dict = {
"name": name,
"checkpoint": "juggernaut",
"concept_mode": concept_mode,
"input_images": str(lora_training_urls),
"num_training_images": n_imgs,
"num_augmented_images": len(captions),
"seed": seed,
"resolution": resolution,
"train_batch_size": train_batch_size,
"num_train_epochs": num_train_epochs,
"max_train_steps": max_train_steps,
"is_lora": is_lora,
"prodigy_d_coef": prodigy_d_coef,
"ti_lr": ti_lr,
"ti_weight_decay": ti_weight_decay,
"lora_weight_decay": lora_weight_decay,
"l1_penalty": l1_penalty,
"lora_param_scaler": lora_param_scaler,
"lora_rank": lora_rank,
"snr_gamma": snr_gamma,
"trigger_text": trigger_text,
"segmentation_prompt": segmentation_prompt,
"crop_based_on_salience": crop_based_on_salience,
"use_face_detection_instead": use_face_detection_instead,
"clipseg_temperature": clipseg_temperature,
"left_right_flip_augmentation": left_right_flip_augmentation,
"augment_imgs_up_to_n": augment_imgs_up_to_n,
"checkpointing_steps": checkpointing_steps,
"run_name": run_name,
"hard_pivot": hard_pivot,
"off_ratio_power": off_ratio_power,
"trainig_captions": captions[:50], # avoid sending back too many captions
}
with open(os.path.join(output_dir, "training_args.json"), "w") as f:
json.dump(args_dict, f, indent=4)
train_generator = main(
pretrained_model,
instance_data_dir=os.path.join(input_dir, "captions.csv"),
output_dir=output_dir,
config = TrainingConfig(
name=name,
lora_training_urls=lora_training_urls,
concept_mode=concept_mode,
sd_model_version=sd_model_version,
seed=seed,
resolution=resolution,
train_batch_size=train_batch_size,
num_train_epochs=num_train_epochs,
max_train_steps=max_train_steps,
gradient_accumulation_steps=gradient_accumulation_steps,
l1_penalty=l1_penalty,
prodigy_d_coef=prodigy_d_coef,
checkpointing_steps=10000,
ti_lr=ti_lr,
ti_weight_decay=ti_weight_decay,
snr_gamma=snr_gamma,
lora_weight_decay=lora_weight_decay,
token_dict=token_dict,
inserting_list_tokens=all_token_lists,
verbose=verbose,
checkpointing_steps=checkpointing_steps,
scale_lr=False,
allow_tf32=True,
mixed_precision="bf16",
#mixed_precision="fp16", # this 100% breaks training... Figure out why!!?
device="cuda:0",
unet_lr=unet_lr,
lora_rank=lora_rank,
is_lora=is_lora,
args_dict=args_dict,
use_dora=use_dora,
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
debug=debug,
hard_pivot=hard_pivot,
off_ratio_power=off_ratio_power,
)
train_generator = train(config=config)
print(f"Debug: {debug}")
while True:
try:
@@ -403,33 +146,9 @@ class Predictor(BasePredictor):
if not debug:
yield CogOutput(name=name, progress=np.round(progress_f, 2))
except StopIteration as e:
output_save_dir, validation_prompts = e.value # Capture the return value
config, output_save_dir = e.value # Capture the return value
break
if not debug:
keys_to_keep = [
"name",
"checkpoint",
"concept_mode",
"input_images",
"num_training_images",
"seed",
"resolution",
"max_train_steps",
"lora_rank",
"trigger_text",
"left_right_flip_augmentation",
"run_name",
"trainig_captions"]
args_dict = {k: v for k, v in args_dict.items() if k in keys_to_keep}
args_dict["grid_prompts"] = validation_prompts
# save final training_args:
final_args_dict_path = os.path.join(output_dir, "training_args.json")
with open(final_args_dict_path, "w") as f:
json.dump(args_dict, f, indent=4)
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
out_path = f"{clean_filename(name)}_eden_concept_lora_{int(time.time())}.tar"
directory = cogPath(output_save_dir)
@@ -443,18 +162,18 @@ class Predictor(BasePredictor):
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
attributes = {}
attributes['grid_prompts'] = validation_prompts
runtime = time.time() - start_time
attributes['job_time_seconds'] = runtime
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
print(f"LORA training finished in {runtime:.1f} seconds")
print(f"LORA training finished in {config.job_time:.1f} seconds")
print(f"Returning {out_path}")
if DEBUG_MODE or debug:
yield cogPath(out_path)
else:
# clear the output_directory to avoid running out of space on the machine:
#shutil.rmtree(output_dir)
yield CogOutput(files=[cogPath(out_path)], name=name, thumbnails=[cogPath(validation_grid_img_path)], attributes=args_dict, isFinal=True, progress=1.0)
yield CogOutput(files=[cogPath(out_path)], name=name, thumbnails=[cogPath(validation_grid_img_path)], attributes=config.dict(), isFinal=True, progress=1.0)
+22
View File
@@ -0,0 +1,22 @@
torch==2.1.0
torchaudio==2.1.0
torchvision==0.16.0
transformers==4.38.0
diffusers==0.26.0
tokenizers==0.15.2
huggingface-hub==0.22.2
ujson==5.10.0
scipy==1.14.0
peft==0.10.0
invisible-watermark==0.2.0
pandas==2.2.1
numpy==1.26.4
opencv-python==4.10.0.84
mediapipe==0.10.14
openai==1.35.13
python-dotenv==1.0.1
prodigyopt==1.0
omegaconf==2.3.0
ujson==5.10.0
bitsandbytes==0.43.1
setuptools==70.3.0
+161
View File
@@ -0,0 +1,161 @@
"""
Faces:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
Objects:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/koji_color.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/plantoid_imgs.zip
Styles:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
"""
import random, os, ast, json, shutil
from itertools import product
import time
from tqdm import tqdm
random.seed(int(1000*time.time()))
def hamming_distance(dict1, dict2):
distance = 0
for key in dict1.keys():
if dict1[key] != dict2.get(key, None):
distance += 1
return distance
#######################################################################################
# Setup the base experiment config:
exp_name = "beeple"
caption_prefix = ""
mask_target_prompts = ""
n_exp = 200 # how many random experiment settings to generate
min_hamming_distance = 1 # min_n_params that have to be different from any previous experiment to be scheduled
nohup = True
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values
# The params are sampled stochastically, so if you want to use a specific value more often, just put it in multiple times
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sdxl"],
"lora_training_urls": [
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.8],
"disable_ti": ['false', 'true'],
"seed": [0],
"resolution": [512],
"train_batch_size": [4],
"n_sample_imgs": [8],
"max_train_steps": [1200],
"checkpointing_steps": [200],
"gradient_accumulation_steps": [1],
"n_tokens": [2],
"ti_lr": [0.001],
"ti_weight_decay": [0.001],
"l1_penalty": [0.0],
"token_warmup_steps": [0],
"tok_cov_reg_w": [2000],
"unet_lr": [0.0002, 0.00005],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16],
"use_dora": ['false'],
"unet_optimizer_type": ['AdamW8bit'],
"is_lora": ['false'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["blip", "gpt4-v"],
"augment_imgs_up_to_n": [40],
"verbose": ['true'],
"debug": ['true']
}
#######################################################################################
# Create a set to hold the combinations that have already been run
scheduled_experiments = set()
# if config_output_dir exists, remove it:
config_output_dir = f"gridsearch_configs/{exp_name}"
shutil.rmtree(config_output_dir, ignore_errors=True)
os.makedirs(config_output_dir, exist_ok=True)
# Open the shell script file
try_sampling_n_times = 200
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
resamples, combination = 0, None
while resamples < try_sampling_n_times:
experiment_settings = {name: random.choice(values) for name, values in hyperparameters.items()}
resamples += 1
min_distance = float('inf')
for str_experiment_settings in scheduled_experiments:
existing_experiment_settings = dict(sorted(ast.literal_eval(str_experiment_settings)))
distance = hamming_distance(experiment_settings, existing_experiment_settings)
min_distance = min(min_distance, distance)
if min_distance >= min_hamming_distance:
str_experiment_settings = str(sorted(experiment_settings.items()))
scheduled_experiments.add(str_experiment_settings)
# Save the experiment to a JSON file
config_filename = f"{config_output_dir}/{exp_name}_{exp_index:03d}.json"
dirname = os.path.dirname(config_filename)
os.makedirs(dirname, exist_ok=True)
# Make some final adjustments to the experiment settings before saving to disk:
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
with open(config_filename, "w") as f:
json.dump(experiment_settings, f, indent=4)
break
if resamples >= try_sampling_n_times:
print(f"\nCould not find a new experiment_setting after random sampling {try_sampling_n_times} times, dumping all experiment_settings to .json files")
break
print(f"\n\n---> Saved {len(scheduled_experiments)} experiment configurations to {config_output_dir}")
def generate_sh_script(folder_path, output_sh_path):
# Get a list of JSON files in the folder, sorted alphabetically
json_files = sorted([f for f in os.listdir(folder_path) if f.endswith('.json')])
# Open the output .sh file for writing
with open(output_sh_path, 'w') as sh_file:
# Write the shebang line for a bash script
sh_file.write("#!/bin/bash\n\n")
# Write a command for each JSON file
for json_file in json_files:
file_path = os.path.join("scripts/", folder_path, json_file)
command = f"python main.py {file_path}\n"
if nohup:
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path)
print(f"\n---> Saved the executable shell script to {output_sh_path}")
+236
View File
@@ -0,0 +1,236 @@
import argparse
from trainer.inference import render_images_eval
from trainer.utils.json_stuff import save_as_json
from trainer.config import TrainingConfig
from trainer.models import pretrained_models
from trainer.utils.io import download
import clip
from PIL import Image
import torch
import numpy as np
import os
from creator_lora.models.resnet50 import ResNet50MLP
"""
todos:
- run eval on user-defined captions
"""
aesthetic_model_checkpoint_filename = "aesthetic_score_best_model.pth"
device = "cuda" if torch.cuda.is_available() else "cpu"
def get_filenames_in_a_folder(folder: str):
"""
returns the list of paths to all the files in a given folder
"""
if folder[-1] == '/':
folder = folder[:-1]
files = os.listdir(folder)
files = [f'{folder}/' + x for x in files]
return files
def get_all_jpg_filenames(folder):
all_filenames = get_filenames_in_a_folder(folder=folder)
jpg_filenames = [filename for filename in all_filenames if filename.lower().endswith('.jpg')]
assert len(jpg_filenames)>0, f"Expected to find at least 1 jpg file but got 0"
return jpg_filenames
def filter_prompt(prompt, remove_this = "in the style of <s0><s1>,", replace_with = ""):
assert remove_this in prompt, f"Expected '{remove_this}' to be present in the prompt: '{prompt}'"
return prompt.replace(
remove_this,
replace_with
)
def get_similarity_matrix(a, b, eps=1e-8):
"""
finds the cosine similarity matrix between each item of a w.r.t each item of b
a and b are expected to be 2 dimensional
added eps for numerical stability
source: https://stackoverflow.com/a/58144658
"""
a_n, b_n = a.norm(dim=1)[:, None], b.norm(dim=1)[:, None]
a_norm = a / torch.max(a_n, eps * torch.ones_like(a_n))
b_norm = b / torch.max(b_n, eps * torch.ones_like(b_n))
sim_mt = torch.mm(a_norm, b_norm.transpose(0, 1))
return sim_mt
class Evaluation:
def __init__(self, image_filenames: list):
self.image_filenames = image_filenames
self.image_features = None
def obtain_image_features(self):
if self.image_features is None:
all_image_features = []
model, preprocess = clip.load("ViT-B/32", device=device)
for f in self.image_filenames:
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
with torch.no_grad():
image_features = model.encode_image(image)
all_image_features.append(image_features.float())
all_image_features = torch.cat(all_image_features, dim = 0)
self.image_features = all_image_features
return self.image_features
def obtain_text_features(self, prompts: list, device):
model, preprocess = clip.load("ViT-B/32", device=device)
text = clip.tokenize(prompts).to(device)
with torch.no_grad():
text_features = model.encode_text(text)
return text_features
def training_image_alignment(self, device, training_image_filenames: list):
generated_image_features = self.obtain_image_features()
training_image_features = []
model, preprocess = clip.load("ViT-B/32", device=device)
for f in training_image_filenames:
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
with torch.no_grad():
image_features = model.encode_image(image)
training_image_features.append(image_features.float())
training_image_features = torch.cat(training_image_features, dim = 0)
return get_similarity_matrix(a=generated_image_features, b=training_image_features).mean().item()
def image_text_alignment(self, device, prompts: list):
image_features = self.obtain_image_features().to(device)
assert image_features.shape[0] == len(prompts), f'Expected len(prompts) ({len(prompts)}) to have the same number of prompts as the number of images provided: {image_features.shape}'
text_features = self.obtain_text_features(prompts=prompts, device=device)
cossim = torch.nn.functional.cosine_similarity(
text_features, image_features, dim = -1
).mean().item()
return cossim
def clip_diversity(self, device: str):
"""
higher = more diverse
"""
all_image_features = self.obtain_image_features().to(device)
distances = 1 - get_similarity_matrix(all_image_features, all_image_features)
assert distances.shape == (
all_image_features.shape[0],
all_image_features.shape[0]
), f'Expected the shape of the distance matrix to be (num_images, num_images) i.e {(all_image_features.shape[0], all_image_features.shape[0])} but got: {distances.shape}'
distances = distances.detach().cpu().numpy()
# Get the upper triangle:
upper_triangle = np.triu(distances, k=1).flatten()
return upper_triangle.mean().item()
def aesthetic_score(self, device: str, checkpoint_path: str):
# assert os.path.exists(checkpoint_path), f"invalid checkpoint_path: {checkpoint_path}"
model = ResNet50MLP(
model_path=checkpoint_path,
device = device
)
scores = []
for f in self.image_filenames:
score = model.predict_score(pil_image=Image.open(f))
scores.append(score)
return sum(scores)/len(scores)
def parse_arguments():
parser = argparse.ArgumentParser(description="Script for generating images based on prompts and computing similarities.")
parser.add_argument("--config_filename", type=str, required=True, default = "sdxl", help="path to config json file")
parser.add_argument("--checkpoint_folder", type=str, required=True,
help="Path to folder containing the checkpoint. Usually a folder which is named like: .../checkpoint-500")
parser.add_argument("--output_json", type=str, required=True,
help="Path to json where we save result values")
parser.add_argument("--output_folder", type=str, required=True,
help="style or face")
parser.add_argument("--training_images_folder", type=str, required=True,
help="path to folder containing training image jpg files. Usually the `images_in` folder")
args = parser.parse_args()
## validate args
assert os.path.exists(args.checkpoint_folder), f"Invalid lora_path: {args.checkpoint_folder}"
assert os.path.exists(args.config_filename), f"Invalid lora_path: {args.config_filename}"
assert os.path.exists(args.training_images_folder), f"Invalid training_images_folder: {args.training_images_folder}"
return args
args = parse_arguments()
os.system(f"mkdir -p {args.output_folder}")
if not os.path.exists(aesthetic_model_checkpoint_filename):
download(
url="https://edenartlab-lfs.s3.amazonaws.com/models/aesthetic_score_best_model.pth",
folder="./",
filepath=None
)
config = TrainingConfig.from_json(args.config_filename)
image_filenames, prompts = render_images_eval(
output_folder=args.output_folder,
concept_mode=config.concept_mode,
render_size=(1024,1024),
checkpoint_folder=args.checkpoint_folder,
pretrained_model=pretrained_models[config.sd_model_version],
seed=0,
is_lora = config.is_lora,
trigger_text='TOK' if config.concept_mode != "style" else ", in the style of TOK"
)
print(f"Eval prompts:")
for i, p in enumerate(prompts):
print(f"{i}:{p}")
eval = Evaluation(image_filenames=image_filenames)
clip_diversity = eval.clip_diversity(device=device)
aesthetic_score = eval.aesthetic_score(device=device, checkpoint_path=aesthetic_model_checkpoint_filename)
image_text_alignment = eval.image_text_alignment(device=device, prompts=prompts)
training_image_alignment = eval.training_image_alignment(
device=device,
training_image_filenames=get_all_jpg_filenames(folder=args.training_images_folder)
)
result = {
"sd_model_version": config.sd_model_version,
"checkpoint_folder": os.path.abspath(args.checkpoint_folder),
"concept_mode": config.concept_mode,
"output_folder": args.output_folder,
"training_images_folder":args.training_images_folder,
"scores": {
"clip_diversity": clip_diversity,
"aesthetic_score": aesthetic_score,
"image_text_alignment": image_text_alignment,
"training_image_alignment": training_image_alignment
}
}
save_as_json(
dictionary_or_list=result,
filename=args.output_json
)
print(f"Eval complete. Saved results here: {args.output_json}")
"""
Example command:
python3 evaluate.py \
--output_folder eval_images \
--checkpoint_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0 \
--output_json eval_results_style.json \
--config_filename lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0/training_args.json \
--training_images_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/images_in
"""
+142
View File
@@ -0,0 +1,142 @@
import os
import json
import matplotlib.pyplot as plt
import seaborn as sns
from collections import defaultdict
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.linear_model import LinearRegression
from sklearn.metrics import r2_score
# Define paths
render_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/lora_models"
config_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/xander_adiff_lora"
ignore_threshold_relative = 0.0 # ignore any datapoint with a score below this threshold
filters = {
"resolution": 512
}
output_dir = f"gridsearch_configs/results/{os.path.basename(config_dir)}"
output_suffix = f"{os.path.basename(render_dir)}"
# Initialize a dictionary to hold parameter values and associated scores
parameters = defaultdict(lambda: defaultdict(list))
# Step 1: Loop over each experiment subdirectory
for i, exp_subdir in enumerate(sorted(os.listdir(render_dir))):
exp_path = os.path.join(render_dir, exp_subdir)
checkpoints_path = os.path.join(exp_path, "checkpoints")
# Step 2: Get the score by counting the number of .jpg files in the checkpoints subdir
if os.path.isdir(checkpoints_path):
score = sum(1 for _ in os.listdir(checkpoints_path) if _.endswith('.jpg'))
# Match the experiment folder with its corresponding JSON file
json_file_name = exp_subdir.split('--')[0] + ".json"
json_file_name = json_file_name.replace('__','_')
json_path = os.path.join(config_dir, json_file_name)
# Step 3: Load the corresponding .json file
if os.path.isfile(json_path):
with open(json_path, 'r') as file:
config = json.load(file)
# Filter out experiments that do not match the filters
if not all(config[key] == value for key, value in filters.items()):
continue
# Step 4: Append all key/value pairs to the total experiment dictionary
for key, value in config.items():
parameters[key]['values'].append(value)
parameters[key]['scores'].append(score)
else:
print(f"Could not find JSON file for experiment {exp_subdir}")
# Print the parameters['output_dir'] with the highest scores (there are usually multiple ties):
max_score = max(parameters['output_dir']['scores'])
best_output_dirs = [output_dir for output_dir, score in zip(parameters['output_dir']['values'], parameters['output_dir']['scores']) if score == max_score]
for best_output_dir in best_output_dirs:
print(f"Best output_dir: {best_output_dir} with score {max_score}")
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.linear_model import LinearRegression
from sklearn.metrics import r2_score
from sklearn.preprocessing import LabelEncoder
os.makedirs(output_dir, exist_ok=True)
print(f"Saving results to {output_dir}...")
def plot_parameters(parameters):
for param, data in parameters.items():
values = np.array(data['values'])
scores = np.array(data['scores'])
# filter based on the ignore_threshold:
ignore_threshold = ignore_threshold_relative * np.max(scores)
mask = scores > ignore_threshold
values = values[mask]
scores = scores[mask]
noise_strength_values = 0.02
noise_strength_scores = 0.02
# Initialize variables for original categorical labels
original_labels = None
# Determine if values are numeric
if values.dtype.kind in 'bifc': # Numeric types
# Add noise directly to values
jittered_values = values + np.random.normal(0, noise_strength_values * (np.max(values) - np.min(values)), values.shape)
else:
# Encode string values to integers for plotting
encoder = LabelEncoder()
original_labels = values.copy()
values = encoder.fit_transform(values)
jittered_values = values + np.random.normal(0, 0.1, values.shape)
# Skip plotting if there is only one unique value for the parameter
if len(np.unique(values)) <= 1:
continue
# Fit a linear regression model to the encoded values if categorical
model = LinearRegression()
values_reshaped = values.reshape(-1, 1) # Reshape for sklearn
model.fit(values_reshaped, scores)
predicted_scores = model.predict(values_reshaped)
# add some jitter to the scores:
jittered_scores = scores + np.random.normal(0, noise_strength_scores * np.max(scores), scores.shape)
# Calculate R² value
r_squared = r2_score(scores, predicted_scores)
# Plot data points
sns.scatterplot(x=jittered_values, y=jittered_scores, alpha=0.6)
# Plot trendline
sns.lineplot(x=np.sort(values), y=predicted_scores[np.argsort(values)], color='red', label=f'R²={r_squared:.2f}')
# Set plot title and labels
plt.title(f'Influence of {param} on the score')
if original_labels is not None:
# Set x-axis labels to the original categorical labels
unique_values = np.unique(values)
plt.xticks(ticks=unique_values, labels=encoder.inverse_transform(unique_values), rotation=45, ha='right')
else:
plt.xlabel(param)
plt.ylabel('Score')
plt.legend()
# Save and close the plot
plt.savefig(f'{output_dir}/res_{param}_{output_suffix}.png')
plt.close()
# Call the updated function with your parameters dictionary
plot_parameters(parameters)
+92
View File
@@ -0,0 +1,92 @@
from diffusers import DDPMScheduler, EulerDiscreteScheduler, StableDiffusionPipeline, StableDiffusionXLPipeline
from peft import PeftModel
import numpy as np
import torch
from huggingface_hub import hf_hub_download
import os, json, random, time, sys
sys.path.append('.')
sys.path.append('..')
from trainer.models import load_models, pretrained_models
from trainer.utils.val_prompts import val_prompts
from trainer.utils.io import make_validation_img_grid
from trainer.utils.utils import seed_everything, pick_best_gpu_id
from trainer.inference import encode_prompt_advanced
from trainer.checkpoint import load_checkpoint
if __name__ == "__main__":
model_version = "sd15"
lora_path = 'lora_models/XANDER_SD15_SWEEP/sd15_face_sweep__004--29_20-43-17-sd15_face_dora_640_1.0_blip_800/checkpoints/checkpoint-800'
lora_scales = np.linspace(0.6, 0.9, 4)
token_scale = None # None means it well get automatically set using lora_scale
render_size = (576, 704) # H,W
n_imgs = 14
n_loops = 2
n_steps = 35
guidance_scale = 7.5
seed = 12
use_lightning = 0
#####################################################################################
pretrained_model = pretrained_models[model_version]
output_dir = f'rendered_images/{lora_path.split("/")[-1]}'
os.makedirs(output_dir, exist_ok=True)
seed_everything(seed)
pick_best_gpu_id()
pipe = load_checkpoint(
pretrained_model_version=model_version,
pretrained_model_path=pretrained_model["path"],
checkpoint_folder=lora_path,
is_lora=True,
device="cuda:0"
)
if use_lightning:
repo = "ByteDance/SDXL-Lightning"
ckpt = "sdxl_lightning_8step_lora.safetensors" # Use the correct ckpt for your step setting!
pipe.load_lora_weights(hf_hub_download(repo, ckpt))
pipe.fuse_lora()
n_steps = 8
guidance_scale=1.5
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
if training_args["concept_mode"] == "style":
validation_prompts_raw = random.choices(val_prompts['style'], k=n_imgs)
elif training_args["concept_mode"] == "face":
validation_prompts_raw = random.choices(val_prompts['face'], k=n_imgs)
else:
validation_prompts_raw = random.choices(val_prompts['object'], k=n_imgs)
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
pipeline_args = {
"num_inference_steps": n_steps,
"guidance_scale": guidance_scale,
"height": render_size[0],
"width": render_size[1],
}
for jj in range(n_loops):
for i in range(len(validation_prompts_raw)):
for lora_scale in lora_scales:
seed += 1
pipe = set_adapter_scales(pipe, lora_scale=lora_scale)
generator = torch.Generator(device='cuda').manual_seed(seed)
c, uc, pc, puc = encode_prompt_advanced(pipe, lora_path, validation_prompts_raw[i], negative_prompt, lora_scale, guidance_scale, concept_mode = training_args["concept_mode"], token_scale = token_scale)
pipeline_args['prompt_embeds'] = c
pipeline_args['negative_prompt_embeds'] = uc
if pretrained_model['version'] == 'sdxl':
pipeline_args['pooled_prompt_embeds'] = pc
pipeline_args['negative_pooled_prompt_embeds'] = puc
image = pipe(**pipeline_args, generator=generator).images[0]
image.save(os.path.join(output_dir, f"{validation_prompts_raw[i][:40]}_seed_{seed}_{i}_lora_scale_{lora_scale:.2f}_{int(time.time())}.jpg"), format="JPEG", quality=95)
seed += 1
-26
View File
@@ -1,26 +0,0 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=3"
cog predict --gpus $GPU_ID \
-i run_name="clipx_sdxl" \
-i caption_prefix="in the style of TOK, " \
-i concept_mode="style" \
-i train_batch_size="4" \
-i sd_model_version="sdxl" \
-i max_train_steps="500" \
-i checkpointing_steps="125" \
-i debug="True" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
-i seed="0"
cog predict --gpus $GPU_ID \
-i run_name="clipx_sd15" \
-i caption_prefix="in the style of TOK, " \
-i concept_mode="style" \
-i train_batch_size="4" \
-i sd_model_version="sd15" \
-i max_train_steps="500" \
-i checkpointing_steps="125" \
-i debug="True" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
-i seed="0"
+8
View File
@@ -0,0 +1,8 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=0"
python main.py train_configs/training_args_face_sdxl.json
python main.py train_configs/training_args_face_sd15.json
python main.py train_configs/training_args_object.json
python main.py train_configs/training_args_style_sd15.json
python main.py train_configs/training_args_style_sdxl.json
+27
View File
@@ -0,0 +1,27 @@
{
"name": "xander_test",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 4,
"max_train_steps": 200,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+28
View File
@@ -0,0 +1,28 @@
{
"name": "xander_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"n_tokens": 2,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.00,
"lora_rank": 4,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -0,0 +1,27 @@
{
"name": "xander_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 300,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -0,0 +1,27 @@
{
"name": "xander_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+28
View File
@@ -0,0 +1,28 @@
{
"name": "banny_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 800,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -0,0 +1,27 @@
{
"name": "clipx_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -0,0 +1,28 @@
{
"name": "clipx_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
-2
View File
@@ -1,2 +0,0 @@
from .trainer import Trainer
from .config import TrainerConfig
+276
View File
@@ -0,0 +1,276 @@
import os, json
from peft.utils import get_peft_model_state_dict
from diffusers.utils import (
convert_all_state_dict_to_peft,
convert_state_dict_to_diffusers,
convert_state_dict_to_kohya,
convert_unet_state_dict_to_peft,
)
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
from safetensors.torch import load_file, save_file
import torch
from diffusers import EulerDiscreteScheduler
from peft import PeftModel
from .utils.json_stuff import save_as_json
from typing import Dict
from trainer.embedding_handler import TokenEmbeddingsHandler
def load_ti_embeddings(pipe, save_path):
# Load the textual_inversion token embeddings into the pipeline:
try: #SDXL
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
except: #SD15
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
embeddings_path = [f for f in os.listdir(save_path) if f.endswith("embeddings.safetensors")][0]
print(f"Loading pretrained token embeddings from {embeddings_path}")
handler.load_embeddings(os.path.join(save_path, embeddings_path))
def set_adapter_scales(pipe, lora_scale = 1.0):
"""
update the pipe with the lora model and the token embeddings
"""
# this loads the lora model into the pipeline at full strength (1.0)
#pipe.unet.load_adapter(lora_path, "eden_lora")
#peft_model.set_adapter(["adapter1", "adapter2"]) # activate both adapters
# First lets see if any lora's are active and unload them:
#pipe.unet.unmerge_adapter()
list_adapters_component_wise = pipe.get_list_adapters()
print(f"list_adapters_component_wise: {list_adapters_component_wise}")
if 1:
for key in list_adapters_component_wise:
adapter_names = list_adapters_component_wise[key]
for adapter_name in adapter_names:
print(f"Set adapter '{adapter_name}' of '{key}' with scale = {lora_scale:.2f}")
pipe.set_adapters(adapter_name, adapter_weights=[lora_scale])
#pipe.unet.merge_adapter()
return pipe
def remove_delimiter_characters(name: str):
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
# Convert to WebUI format
def convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename: str,
output_filename: str
):
assert os.path.exists(pytorch_lora_weights_filename), f"Invalid path: {pytorch_lora_weights_filename}"
lora_state_dict = load_file(pytorch_lora_weights_filename)
peft_state_dict = convert_all_state_dict_to_peft(lora_state_dict)
kohya_state_dict = convert_state_dict_to_kohya(peft_state_dict)
# This is a very custom hack because for some reason these 'base_model_model_' prefixes are added to the keys and ComfyUI does not like them...
replace_dict = {"base_model_model_": ""}
# enumerate and apply replace_dict:
for key in list(kohya_state_dict.keys()):
for old_key, new_key in replace_dict.items():
if old_key in key:
new_key = key.replace(old_key, new_key)
kohya_state_dict[new_key] = kohya_state_dict.pop(key)
save_file(kohya_state_dict, output_filename)
def save_checkpoint(
output_dir: str,
global_step: int,
unet,
embedding_handler,
token_dict: dict,
is_lora: bool,
unet_lora_parameters,
pretrained_model_version: str,
name: str = None,
text_encoder_peft_models: list = [None]
):
"""
Save the model's embeddings and special parameters (Lora) to the specified directory.
Note: This function directly corresponds to the `load_checkpoint` method
Args:
`output_dir` (str): The directory path where the checkpoint will be saved.
`global_step` (int): The current global step or epoch number.
`unet`: The main model to save.
`embedding_handler`: The handler for saving embeddings.
`token_dict` (dict): Special parameters associated with the model.
`is_lora` (bool): Whether the model includes LoRA components.
`unet_lora_parameters`: Parameters associated with the LoRA components.
`name` (str, optional): Name identifier for the checkpoint. Defaults to None.
`text_encoder_peft_models` (list, optional): List of additional text encoder models to save. Defaults to None.
Returns:
None
Saves:
- {name}_embeddings.safetensors: Embeddings of the model.
- special_params.json: Special parameters of the model.
If `text_encoder_peft_models` is provided, saves each model in a separate directory with the
following structure:
- text_encoder_lora_{index}/
- adapter_config.json
- adapter_model.safetensors
- README.md
If `is_lora` is True, saves additional LoRA-related data:
- LoRA weights
- LoRA weights converted for web UI
If `is_lora` is False then it assumes that it's a vanilla unet model and saves it in the usual huggingface way.
"""
print(f"Saving checkpoint at step.. {global_step}")
name = remove_delimiter_characters(name)
embedding_handler.save_embeddings(
os.path.join(
output_dir,
f"{name}_{pretrained_model_version}_embeddings.safetensors"
)
)
save_as_json(
token_dict,
filename = os.path.join(
output_dir, "special_params.json"
)
)
if is_lora:
assert len(unet_lora_parameters) > 0, f"Expected len(unet_lora_parameters) to be greater than zero if is_lora is True"
# This saves adapter_config.json:
# TODO: adjust inference.py so it can load everything without needing this file
unet.save_pretrained(save_directory = output_dir)
text_encoder_lora_layers = [None, None]
for idx, model in enumerate(text_encoder_peft_models):
if model is not None:
lora_tensors = get_peft_model_state_dict(model)
text_encoder_lora_layers[idx] = convert_state_dict_to_diffusers(lora_tensors)
lora_tensors = get_peft_model_state_dict(unet)
unet_lora_layers_to_save = convert_state_dict_to_diffusers(lora_tensors)
if pretrained_model_version == "sdxl":
print("Saving LoRA weights for SDXL model...")
StableDiffusionXLPipeline.save_lora_weights(
output_dir,
unet_lora_layers=unet_lora_layers_to_save,
text_encoder_lora_layers=text_encoder_lora_layers[0],
text_encoder_2_lora_layers=text_encoder_lora_layers[1],
)
elif pretrained_model_version == "sd15":
print("Saving LoRA weights for SD15 model...")
StableDiffusionPipeline.save_lora_weights(
output_dir,
unet_lora_layers=unet_lora_layers_to_save,
text_encoder_lora_layers=text_encoder_lora_layers[0],
)
else:
raise ValueError(
f"Invalid pretrained_model_version: {pretrained_model_version}. Expected one of: 'sdxl' or 'sd15'"
)
convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_LoRa.safetensors")
)
else:
# Save the entire, finetuned unet weights:
unet.save_pretrained(save_directory = output_dir)
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
return
def load_checkpoint(
pretrained_model_version: str,
pretrained_model_path: str,
lora_save_path: str,
is_lora: bool,
device: str,
lora_scale: float = 1.0,
):
"""
Load a pre-trained model checkpoint and prepare it for inference.
Note: This function directly corresponds to the `save_checkpoint` method
Args:
`pretrained_model_version` (`str`): Version of the pre-trained model (`sd15` or `sdxl`).
`pretrained_model_path` (`str`): Path to the pre-trained model file.
`lora_save_path` (`str`): Path to the LoRa checkpoint folder.
`is_lora` (`bool`): Whether LoRA model components are used.
`device` (Union[`str`, `torch.device`]): Device for inference.
Raises:
NotImplementedError: If an unsupported `pretrained_model_version` is provided.
"""
assert os.path.exists(pretrained_model_path), f"Invalid pretrained_model_path: {pretrained_model_path}"
if pretrained_model_version == "sd15":
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model_path, torch_dtype=torch.float16, use_safetensors=True)
elif pretrained_model_version == "sdxl":
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model_path, torch_dtype=torch.float16, use_safetensors=True)
else:
raise NotImplementedError(f"Invalid pretrained_model_version: {pretrained_model_version}")
pipe = pipe.to(device, dtype=torch.float16)
print(f"Loaded new {pretrained_model_version} model from: {pretrained_model_path}")
# Load textual_inversion embeddings:
load_ti_embeddings(pipe, lora_save_path)
# TODO: why does this give key errors???
#pipe.load_lora_weights(lora_save_path, weight_name='pytorch_lora_weights.safetensors')
#pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
#pipe.fuse_lora(lora_scale=lora_scale)
#return pipe
assert os.path.exists(lora_save_path), f"Invalid lora_save_path: {lora_save_path}"
text_encoder_0_path = os.path.join(
lora_save_path, "text_encoder_lora_0"
)
text_encoder_1_path = os.path.join(
lora_save_path, "text_encoder_lora_1"
)
if os.path.exists(
text_encoder_0_path
):
pipe.text_encoder = PeftModel.from_pretrained(pipe.text_encoder, text_encoder_0_path)
print(f"loaded text_encoder LoRA from: {text_encoder_0_path}")
if os.path.exists(
text_encoder_1_path
):
pipe.text_encoder_2 = PeftModel.from_pretrained(pipe.text_encoder_2, text_encoder_1_path)
print(f"loaded text_encoder LoRA from: {text_encoder_1_path}")
if is_lora:
pipe.unet = PeftModel.from_pretrained(model = pipe.unet, model_id = lora_save_path)
else:
pipe.unet = pipe.unet.from_pretrained(lora_save_path)
print(f"Successfully loaded full checkpoint for inference!")
pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
return pipe
+169 -48
View File
@@ -1,61 +1,182 @@
from typing import Optional, List, Dict, Any
from pydantic import BaseModel, Field
import random
import json
from typing import Union, List, Optional
from datetime import datetime
from pydantic import BaseModel
import json, time, os
from typing import Literal
import torch
from trainer.utils.utils import pick_best_gpu_id
precision_map = {
"fp16": torch.float16,
"bf16": torch.bfloat16,
"fp32": torch.float32
class ModelPaths:
def __init__(self):
self.paths = {
"BLIP": "./cache",
"CLIP": "./cache",
"SR": "./cache",
"SD": "./models",
}
def get_path(self, key):
return self.paths.get(key, None)
def set_path(self, key, path):
if key in self.paths:
self.paths[key] = path
model_paths = ModelPaths()
# Default download urls in case no local model is found:
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-XL-v6/resolve/main/juggernautXL_version6Rundiffusion.safetensors"
SDXL_URL = "https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_base_1.0_0.9vae.safetensors"
SD15_URL = "https://huggingface.co/KamCastle/jugg/resolve/main/juggernaut_reborn.safetensors"
pretrained_models = {
"sdxl": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SDXL_URL)), "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SD15_URL)), "url": SD15_URL, "version": "sd15"}
}
class TrainerConfig(BaseModel, extra = "forbid"):
pretrained_model: Dict[str, str] # should be a dict with keys "path" and "version"
name: str='unnamed',
trigger_text: str='a photo of TOK, ',
instance_data_dir: str = "./dataset/zeke/captions.csv"
concept_mode: Literal["face", "concept", "object", "style"]
output_dir: str = "lora_output"
seed: Optional[int] = Field(default_factory=lambda: random.randint(0, 2**32 - 1))
resolution: int = 960
class TrainingConfig(BaseModel):
lora_training_urls: str
concept_mode: Literal["face", "style", "object"]
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
caption_model: Literal["gpt4-v", "blip"] = "blip"
sd_model_version: Literal["sdxl", "sd15", None] = None
ckpt_path: str = None # optional hardcoded checkpoint path
pretrained_model: dict = None
seed: Union[int, None] = None
resolution: int = 512
validation_img_size: Optional[Union[int, List[int]]] = None # [width, height], target_n_pixels ** 0.5 or None
train_img_size: List[int] = None
train_aspect_ratio: float = None
train_batch_size: int = 4
num_train_epochs: int = 10000
max_train_steps: int = 360
checkpointing_steps: int = 10000
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 1.0e-3
prodigy_d_coef: float = 1.0
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
lora_weight_decay: float = 0.002
ti_lr: float = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done
cond_reg_w: float = 0.0e-5
tok_cond_reg_w: float = 0.0e-5
tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
snr_gamma: float = 5.0
lora_alpha_multiplier: float = 1.0
lora_rank: int = 12
use_dora: bool = False
left_right_flip_augmentation: bool = True
augment_imgs_up_to_n: int = 40
mask_target_prompts: Union[None, str] = None
crop_based_on_salience: bool = True
use_face_detection_instead: bool = False # use a different model (not CLIPSeg) to generate face masks
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
n_sample_imgs: int = 4
name: str = None
output_dir: str = "eden_lora_training_runs"
debug: bool = False
allow_tf32: bool = True
disable_ti: bool = False
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 2
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
token_dict: dict = {"TOK": "<s0><s1>"}
device: str = "cuda:0"
crops_coords_top_left_h: int = 0
crops_coords_top_left_w: int = 0
train_batch_size: int = 1
train_dataset_cache: bool = True
num_train_epochs: int = 10000
max_train_steps: Optional[int] = None
checkpointing_steps: int = 500000
gradient_accumulation_steps: int = 1
do_cache: bool = True
unet_learning_rate: float = 1.0
textual_inversion_lr: float = 1e-3
textual_inversion_weight_decay: float = 3e-4
prodigy_d_coef: float = 0.5,
l1_penalty: float = 0.0
lora_weight_decay: float = 0.005
scale_lr_based_on_grad_acc: bool = False
lr_scheduler_name: str = "constant"
lr_warmup_steps: int = 50
lr_num_cycles: int = 1
lr_power: float = 1.0
snr_gamma: float = 5.0
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0
allow_tf32: bool = True
precision: Literal["bf16", "fp16", "fp32"] = "bf16"
optimizer_name: Literal["prodigy", "adamw"] = "prodigy"
device: str = "cuda"
token_dict: Dict[str, str] = {"TOK": "<s0><s1>"}
inserting_list_tokens: List[str] = ["<s0><s1>"]
verbose: bool = True
is_lora: bool = True
lora_rank: int = 12
lora_alpha: int = 12
args_dict: Dict[str, Any] = {}
debug: bool = False
hard_pivot: bool = True
off_ratio_power: float = 0.1
training_attributes: dict = {}
aspect_ratio_bucketing: bool = False
start_time: float = 0.0
job_time: float = 0.0
"""
For text encoder lora training, the trigger variable is: text_encoder_lora_optimizer
if text_encoder_lora_optimizer is not None then everything else is used.
Else the other variables are ignored.
"""
text_encoder_lora_optimizer: Union[None, Literal["adamw"]] = None
text_encoder_lora_lr: float = 1.0e-5
txt_encoders_lr_warmup_steps: int = 200
text_encoder_lora_weight_decay: float = 1.0e-5
text_encoder_lora_rank: int = 16
def __init__(self, **data):
super().__init__(**data)
if not self.ckpt_path:
self.pretrained_model = pretrained_models[self.sd_model_version]
else:
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
# add some metrics to the foldername:
lora_str = "dora" if self.use_dora else "lora"
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
if not self.name:
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
self.output_dir = self.output_dir + f"/{self.name}/" + f"{timestamp_short}-{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
os.makedirs(self.output_dir, exist_ok=True)
if self.seed is None:
self.seed = int(time.time())
if self.unet_lr_warmup_steps is None:
self.unet_lr_warmup_steps = self.max_train_steps
if self.concept_mode == "face":
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
self.mask_target_prompts = "face"
#self.use_face_detection_instead = True
if not self.sample_imgs_lora_scale:
if self.sd_model_version == "sdxl":
self.sample_imgs_lora_scale = 0.7
else:
self.sample_imgs_lora_scale = 0.85
if self.use_dora:
print(f"Disabling L1 penalty and LoRA weight decay for DORA training.")
self.l1_penalty = 0.0
self.lora_weight_decay = 0.0
self.text_encoder_lora_weight_decay = 0.0
# build the inserting_list_tokens and token dict using n_tokens:
inserting_list_tokens = [f"<s{i}>" for i in range(self.n_tokens)]
self.inserting_list_tokens = inserting_list_tokens
self.token_dict = {"TOK": "".join(inserting_list_tokens)}
gpu_id = pick_best_gpu_id()
self.device = f'cuda:{gpu_id}'
self.start_time = time.time()
@classmethod
def from_json(cls, file_path: str):
with open(file_path, 'r') as f:
data = json.load(f)
return cls(**data)
def save_as_json(self, file_path: str) -> None:
with open(file_path, 'w') as f:
json.dump(self.dict(), f, indent=4)
json.dump(self.dict(), f, indent=4)
+168 -148
View File
@@ -1,171 +1,191 @@
import os
import torch
import numpy as np
import pandas as pd
import PIL
from PIL import Image
from torch.utils.data import Dataset
from typing import Union
from dataclasses import dataclass
import torchvision.transforms as transforms
import numpy as np
import torch
from typing import Tuple, Dict, List
@dataclass
class ImageSize:
width: int
height: int
def prepare_image(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
) -> torch.Tensor:
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
image = pipe.image_processor.preprocess(pil_image)
return image
default_tokenizer_kwargs = dict(
padding="max_length",
max_length=77,
truncation=True,
add_special_tokens=True,
return_tensors="pt"
)
default_image_transforms = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
def convert_pil_mask_to_tensor(mask):
arr = np.array(mask)
def prepare_mask(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
) -> torch.Tensor:
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
arr = np.array(pil_image.convert("L"))
arr = arr.astype(np.float32) / 255.0
return torch.tensor(arr).unsqueeze(0).unsqueeze(0)
class ImageCaptionDataset(Dataset):
def __init__(
self,
image_folder: str,
csv_filename: str,
mask_folder: Union[str, None] = None,
validate_csv: bool = True,
size: Union[ImageSize, None] = None,
):
super().__init__()
assert os.path.exists(image_folder), f"Invalid image_folder: {image_folder}"
assert os.path.exists(csv_filename), f"Invalid csv_filename: {csv_filename}"
df = pd.read_csv(csv_filename)
assert (
"caption" in list(df.columns)
), f"Expected column: 'caption' to exist but got columns: {list(df.columns)}"
assert (
"image_path" in list(df.columns)
), f"Expected column: 'image_path' to exist but got columns: {list(df.columns)}"
self.csv_filename = csv_filename
self.captions = df.caption.values
self.image_path = df.image_path.values
self.size = size
self.image_folder = image_folder
self.mask_folder = mask_folder
if mask_folder is not None:
assert (
"mask_path" in list(df.columns)
), f"Expected column: 'mask_path' to exist but got columns: {list(df.columns)}"
self.mask_path = df.mask_path
else:
self.mask_path = None
if validate_csv:
self.validate_csv()
def validate_csv(self):
for idx in range(len(self.image_path)):
filename = os.path.join(self.image_folder, self.image_path[idx])
assert os.path.exists(
filename
), f"Invalid image path: {filename}\nPlease check your CSV file: {self.csv_filename}"
if self.mask_path is not None:
for idx in range(len(self.mask_path)):
filename = os.path.join(self.image_folder, self.image_path[idx])
assert os.path.exists(
filename
), f"Invalid mask_path: {filename}\nPlease check your CSV file: {self.csv_filename}"
def __getitem__(self, idx: int) -> dict:
filename = os.path.join(self.image_folder, self.image_path[idx])
image = Image.open(filename)
if self.size is not None:
# inspired by dataset_and_utils.py -> prepare_image
image = image.resize(
(self.size.width, self.size.height),
resample=Image.BICUBIC,
reducing_gap=1,
)
image = image.convert("RGB")
caption = self.captions[idx]
if self.mask_path is not None:
image_width, image_height = image.size
mask_filename = os.path.join(self.mask_folder, self.mask_path[idx])
mask = Image.open(mask_filename)
mask = mask.convert("L")
mask = mask.resize(
(image_width, image_height),
resample=Image.BICUBIC,
reducing_gap=1,
)
else:
mask = None
return {"image": image, "caption": caption, "mask": mask}
def __len__(self) -> int:
return len(self.captions)
arr = np.expand_dims(arr, 0)
image = torch.from_numpy(arr).unsqueeze(0)
return image
class PreprocessedDataset(Dataset):
def __init__(
self,
image_caption_dataset: ImageCaptionDataset,
tokenizers: list,
vae,
text_encoders: Union[list, None] = None,
cache: bool = True,
tokenizer_kwargs: dict = default_tokenizer_kwargs,
scale_vae_latents: bool = True
data_dir: str,
pipe,
vae_encoder,
do_cache: bool = False,
size: List[int] = [512, 512],
text_dropout: float = 0.0,
aspect_ratio_bucketing: bool = False,
train_batch_size: int = None, # required for aspect_ratio_bucketing
substitute_caption_map: Dict[str, str] = {},
):
self.image_caption_dataset=image_caption_dataset
self.tokenizers=tokenizers
self.vae=vae
self.text_encoders=text_encoders
self.cache=cache
self.tokenizer_kwargs=tokenizer_kwargs
self.scale_vae_latents=scale_vae_latents
def __getitem__(self, idx: int):
super().__init__()
self.data_dir = data_dir
self.csv_path = os.path.join(data_dir, "captions.csv")
self.data = pd.read_csv(self.csv_path)
data = self.image_caption_dataset[idx]
image, caption, mask = data["image"], data["caption"], data["mask"]
self.captions = self.data["caption"]
self.captions = self.captions.str.lower()
for key, value in substitute_caption_map.items():
self.captions = self.captions.str.replace(key.lower(), value)
tokenized_captions = []
self.image_path = self.data["image_path"]
for tokenizer in self.tokenizers:
tokenized_text = tokenizer(
caption,
**self.tokenizer_kwargs
).input_ids.squeeze()
tokenized_captions.append(
tokenized_text
if "mask_path" not in self.data.columns:
self.mask_path = None
else:
self.mask_path = self.data["mask_path"]
self.pipe = pipe
self.vae_encoder = vae_encoder
self.vae_scaling_factor = self.vae_encoder.config.scaling_factor
self.text_dropout = text_dropout
self.size = size
if do_cache:
print("Caching latents, masks and captions...\n")
self.vae_latents = []
self.masks = []
self.do_cache = True
for idx in range(len(self.data)):
if len(self.data) < 25:
print(self.captions[idx])
vae_latent, mask = self._process(idx)
self.vae_latents.append(vae_latent)
self.masks.append(mask)
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
del self.vae_encoder
else:
self.do_cache = False
if aspect_ratio_bucketing:
print("Using aspect ratio bucketing.")
assert train_batch_size is not None, f"Please also provide a `train_batch_size` when you have set `aspect_ratio_bucketing == True`"
from .utils.aspect_ratio_bucketing import BucketManager
aspect_ratios = {}
for idx in range(len(self.data)):
aspect_ratios[idx] = Image.open(os.path.join(self.data_dir, self.image_path[idx])).size
self.bucket_manager = BucketManager(
aspect_ratios = aspect_ratios,
bsz = train_batch_size,
debug=True
)
else:
print("Not using aspect ratio bucketing.")
self.bucket_manager = None
def get_aspect_ratio_bucketed_batch(self):
assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()"
indices, resolution = self.bucket_manager.get_batch()
print(f"Got bucket batch: {indices}, resolution: {resolution}")
tok1, tok2, vae_latents, masks = [], [], [], []
for idx in indices:
if self.tokenizer_2 is None:
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
else:
(t1, t2), v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
tok2.append(t2.unsqueeze(0))
tok1.append(t1.unsqueeze(0))
vae_latents.append(v.unsqueeze(0))
masks.append(m.unsqueeze(0))
tok1 = torch.cat(tok1, dim = 0)
if self.tokenizer_2 is None:
pass
else:
tok2 = torch.cat(tok2, dim = 0)
vae_latents = torch.cat(vae_latents, dim = 0)
masks = torch.cat(masks, dim = 0)
if self.tokenizer_2 is None:
return (tok1, None), vae_latents, masks
else:
return (tok1, tok2), vae_latents, masks
def __len__(self) -> int:
return len(self.data)
@torch.no_grad()
def _process(
self, idx: int, bucketing_resolution: tuple = None
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
image_path = self.image_path[idx]
image_path = os.path.join(self.data_dir, image_path)
image = PIL.Image.open(image_path).convert("RGB")
if bucketing_resolution is None:
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
else:
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
image_tensor = default_image_transforms(image).unsqueeze(0).to(self.vae.device, dtype = self.vae.dtype)
vae_latent = self.vae_encoder.encode(image).latent_dist
dummy_vae_latent = vae_latent.sample()
if mask is not None:
mask = convert_pil_mask_to_tensor(mask)
if self.mask_path is None:
mask = torch.ones_like(
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
# raise ValueError(image_tensor.mean(), image_tensor.var())
vae_latent = self.vae.encode(image_tensor).latent_dist.sample()
if self.scale_vae_latents:
vae_latent = vae_latent * self.vae.config.scaling_factor
else:
mask_path = self.mask_path[idx]
mask_path = os.path.join(self.data_dir, mask_path)
mask = PIL.Image.open(mask_path)
mask = prepare_mask(mask, self.size[0], self.size[1]).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
mask_dtype = mask.dtype
mask = mask.float()
mask = torch.nn.functional.interpolate(
mask, size=(dummy_vae_latent.shape[-2], dummy_vae_latent.shape[-1]), mode="nearest"
)
mask = mask.to(dtype=mask_dtype)
mask = mask.repeat(1, dummy_vae_latent.shape[1], 1, 1)
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
return vae_latent, mask.squeeze()
def __getitem__(
self, idx: int, bucketing_resolution:tuple = None
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
if self.do_cache:
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
return self.captions[idx], vae_latent.squeeze(), self.masks[idx]
else: # This code pathway has not been tested in a long time and might be broken
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution)
vae_latent = vae_latent.sample() * self.vae_scaling_factor
return caption, vae_latent.squeeze(), mask
return {
"tokenized_captions": tokenized_captions,
"vae_latent": vae_latent.squeeze(),
"mask": mask
}
def __len__(self):
return len(self.image_caption_dataset)
-805
View File
@@ -1,805 +0,0 @@
import os
from typing import Dict, List, Optional, Tuple
import random
import numpy as np
import pandas as pd
import gc
import PIL
import torch
import torch.utils.checkpoint
from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
from PIL import Image
from safetensors import safe_open
from safetensors.torch import save_file
from torch.utils.data import Dataset
from transformers import AutoTokenizer, PretrainedConfig
import torch
import torch.nn.functional as F
import matplotlib.pyplot as plt
def plot_torch_hist(parameters, epoch, checkpoint_dir, name, bins=100, min_val=-1, max_val=1, ymax_f = 0.75):
# Flatten and concatenate all parameters into a single tensor
all_params = torch.cat([p.data.view(-1) for p in parameters])
# Convert to CPU for plotting
all_params_cpu = all_params.cpu().float().numpy()
# Plot histogram
plt.figure()
plt.hist(all_params_cpu, bins=bins, density=False)
plt.ylim(0, ymax_f * len(all_params_cpu))
plt.xlim(min_val, max_val)
plt.xlabel('Weight Value')
plt.ylabel('Count')
plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f}, min = {np.min(all_params_cpu):.2f}, max = {np.max(all_params_cpu):.2f})')
plt.savefig(f"{checkpoint_dir}/{name}_histogram_{epoch:04d}.png")
plt.close()
# plot the learning rates:
def plot_lrs(lora_lrs, ti_lrs, save_path='learning_rates.png'):
plt.figure()
plt.plot(range(len(lora_lrs)), lora_lrs, label='LoRA LR')
plt.plot(range(len(lora_lrs)), ti_lrs, label='TI LR')
plt.yscale('log') # Set y-axis to log scale
plt.ylim(1e-6, 3e-3)
plt.xlabel('Step')
plt.ylabel('Learning Rate')
plt.title('Learning Rate Curves')
plt.legend()
plt.savefig(save_path)
plt.close()
from scipy.signal import savgol_filter
def plot_loss(losses, save_path='losses.png', window_length=31, polyorder=3):
if len(losses) < window_length:
return
smoothed_losses = savgol_filter(losses, window_length, polyorder)
plt.figure()
plt.plot(losses, label='Actual Losses')
plt.plot(smoothed_losses, label='Smoothed Losses', color='red')
# plt.yscale('log') # Uncomment if log scale is desired
plt.xlabel('Step')
plt.ylabel('Training Loss')
plt.legend()
plt.savefig(save_path)
plt.close()
def prepare_image(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
) -> torch.Tensor:
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
arr = np.array(pil_image.convert("RGB"))
arr = arr.astype(np.float32) / 127.5 - 1
arr = np.transpose(arr, [2, 0, 1])
image = torch.from_numpy(arr).unsqueeze(0)
return image
def prepare_mask(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512
) -> torch.Tensor:
pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1)
arr = np.array(pil_image.convert("L"))
arr = arr.astype(np.float32) / 255.0
arr = np.expand_dims(arr, 0)
image = torch.from_numpy(arr).unsqueeze(0)
return image
class PreprocessedDataset(Dataset):
def __init__(
self,
csv_path: str,
tokenizer_1,
tokenizer_2,
vae_encoder,
text_encoder_1=None,
text_encoder_2=None,
do_cache: bool = False,
size: int = 512,
text_dropout: float = 0.0,
scale_vae_latents: bool = True,
substitute_caption_map: Dict[str, str] = {},
):
super().__init__()
self.data = pd.read_csv(csv_path)
self.csv_path = csv_path
self.caption = self.data["caption"]
# make it lowercase
self.caption = self.caption.str.lower()
for key, value in substitute_caption_map.items():
self.caption = self.caption.str.replace(key.lower(), value)
self.image_path = self.data["image_path"]
if "mask_path" not in self.data.columns:
self.mask_path = None
else:
self.mask_path = self.data["mask_path"]
if text_encoder_1 is None:
self.return_text_embeddings = False
else:
self.text_encoder_1 = text_encoder_1
self.text_encoder_2 = text_encoder_2
self.return_text_embeddings = True
assert (
NotImplementedError
), "Preprocessing Text Encoder is not implemented yet"
self.tokenizer_1 = tokenizer_1
self.tokenizer_2 = tokenizer_2
self.vae_encoder = vae_encoder
self.scale_vae_latents = scale_vae_latents
self.text_dropout = text_dropout
self.size = size
if do_cache:
self.vae_latents = []
self.tokens_tuple = []
self.masks = []
self.do_cache = True
print("Captions to train on: ")
for idx in range(len(self.data)):
token, vae_latent, mask = self._process(idx)
self.vae_latents.append(vae_latent)
self.tokens_tuple.append(token)
self.masks.append(mask)
print(f"Cached latents and masks for {len(self.vae_latents)} images.")
del self.vae_encoder
else:
self.do_cache = False
def __len__(self) -> int:
return len(self.data)
@torch.no_grad()
def _process(
self, idx: int
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
image_path = self.image_path[idx]
image_path = os.path.join(os.path.dirname(self.csv_path), image_path)
image = PIL.Image.open(image_path).convert("RGB")
image = prepare_image(image, self.size, self.size).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
caption = self.caption[idx]
print(caption)
# tokenizer_1
ti1 = self.tokenizer_1(
caption,
padding="max_length",
max_length=77,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
).input_ids.squeeze()
if self.tokenizer_2 is None:
ti2 = None
else:
ti2 = self.tokenizer_2(
caption,
padding="max_length",
max_length=77,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
).input_ids.squeeze()
vae_latent = self.vae_encoder.encode(image).latent_dist.sample()
if self.scale_vae_latents:
vae_latent = vae_latent * self.vae_encoder.config.scaling_factor
if self.mask_path is None:
mask = torch.ones_like(
vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
else:
mask_path = self.mask_path[idx]
mask_path = os.path.join(os.path.dirname(self.csv_path), mask_path)
mask = PIL.Image.open(mask_path)
mask = prepare_mask(mask, self.size, self.size).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
mask_dtype = mask.dtype
mask = mask.float()
mask = torch.nn.functional.interpolate(
mask, size=(vae_latent.shape[-2], vae_latent.shape[-1]), mode="nearest"
)
mask = mask.to(dtype=mask_dtype)
mask = mask.repeat(1, vae_latent.shape[1], 1, 1)
assert len(mask.shape) == 4 and len(vae_latent.shape) == 4
if ti2 is None: # sd15
return ti1, vae_latent.squeeze(), mask.squeeze()
else: # sdxl
return (ti1, ti2), vae_latent.squeeze(), mask.squeeze()
def atidx(
self, idx: int
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
if self.do_cache:
return self.tokens_tuple[idx], self.vae_latents[idx], self.masks[idx]
else:
return self._process(idx)
def __getitem__(
self, idx: int
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]:
token, vae_latent, mask = self.atidx(idx)
return token, vae_latent, mask
def import_model_class_from_model_name_or_path(
pretrained_model_name_or_path: str, revision: str, subfolder: str = "text_encoder"
):
text_encoder_config = PretrainedConfig.from_pretrained(
pretrained_model_name_or_path, subfolder=subfolder, revision=revision
)
model_class = text_encoder_config.architectures[0]
if model_class == "CLIPTextModel":
from transformers import CLIPTextModel
print("Importing CLIPTextModel")
return CLIPTextModel
elif model_class == "CLIPTextModelWithProjection":
from transformers import CLIPTextModelWithProjection
print("Importing CLIPTextModelWithProjection")
return CLIPTextModelWithProjection
else:
raise ValueError(f"{model_class} is not supported.")
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
if not isinstance(pretrained_model, dict) or 'path' not in pretrained_model or 'version' not in pretrained_model:
raise ValueError("pretrained_model must be a dict with 'path' and 'version' keys")
print(f"Loading model weights from {pretrained_model['path']} as dtype: {weight_dtype}...")
if pretrained_model['version'] == "sd15":
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
else:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
pipe = pipe.to(device, dtype=weight_dtype)
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
vae = pipe.vae
unet = pipe.unet
tokenizer_one = pipe.tokenizer
text_encoder_one = pipe.text_encoder
vae.requires_grad_(False)
text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype)
unet.to(device, dtype=weight_dtype)
if keep_vae_float32:
vae.to(device, dtype=torch.float32)
else:
vae.to(device, dtype=weight_dtype)
if weight_dtype != torch.float32:
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but not for training!!")
tokenizer_two = text_encoder_two = None
if pretrained_model['version'] == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
return (
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
)
class TokenEmbeddingsHandler:
def __init__(self, text_encoders, tokenizers):
self.text_encoders = text_encoders
self.tokenizers = tokenizers
self.train_ids: Optional[torch.Tensor] = None
self.inserting_toks: Optional[List[str]] = None
self.embeddings_settings = {}
def get_trainable_embeddings(self):
trainable_embeddings = []
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
trainable_embeddings.append(text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids])
return trainable_embeddings
def find_nearest_tokens(self, query_embedding, tokenizer, text_encoder, idx, distance_metric, top_k = 5):
# given a query embedding, compute the distance to all embeddings in the text encoder
# and return the top_k closest tokens
assert distance_metric in ["l2", "cosine"], "distance_metric should be either 'l2' or 'cosine'"
# get all non-optimized embeddings:
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[index_no_updates]
# compute the distance between the query embedding and all embeddings:
if distance_metric == "l2":
diff = (embeddings - query_embedding.unsqueeze(0))**2
distances = diff.sum(-1)
distances, indices = torch.topk(distances, top_k, dim=0, largest=False)
elif distance_metric == "cosine":
distances = F.cosine_similarity(embeddings, query_embedding.unsqueeze(0), dim=-1)
distances, indices = torch.topk(distances, top_k, dim=0, largest=True)
nearest_tokens = tokenizer.convert_ids_to_tokens(indices)
return nearest_tokens, distances
def print_token_info(self, distance_metric = "cosine"):
print(f"----------- Closest tokens (distance_metric = {distance_metric}) --------------")
current_token_embeddings = self.get_trainable_embeddings()
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
query_embeddings = current_token_embeddings[idx]
for token_id, query_embedding in enumerate(query_embeddings):
nearest_tokens, distances = self.find_nearest_tokens(query_embedding, tokenizer, text_encoder, idx, distance_metric)
# print the results:
print(f"txt-encoder {idx}, token {token_id}: :")
for i, (token, dist) in enumerate(zip(nearest_tokens, distances)):
print(f"---> {distance_metric} of {dist:.4f}: {token}")
idx += 1
def get_start_embedding(self, text_encoder, tokenizer, example_tokens, unk_token_id = 49407, verbose = False, desired_std_multiplier = 0.0):
print('-----------------------------------------------')
# do some cleanup:
example_tokens = [tok.lower() for tok in example_tokens]
example_tokens = list(set(example_tokens))
starting_ids = tokenizer.convert_tokens_to_ids(example_tokens)
# filter out any tokens that are mapped to unk_token_id:
example_tokens = [tok for tok, tok_id in zip(example_tokens, starting_ids) if tok_id != unk_token_id]
starting_ids = [tok_id for tok_id in starting_ids if tok_id != unk_token_id]
if verbose:
print("Token mapping:")
for i, token in enumerate(example_tokens):
print(f"{token} -> {starting_ids[i]}")
embeddings, stds = [], []
for i, token_index in enumerate(starting_ids):
embedding = text_encoder.text_model.embeddings.token_embedding.weight.data[token_index].clone()
embeddings.append(embedding)
stds.append(embedding.std())
#print(f"token: {example_tokens[i]}, embedding-std: {embedding.std():.4f}, embedding-mean: {embedding.mean():.4f}")
embeddings = torch.stack(embeddings)
#print(f"Embeddings: {embeddings.shape}, std: {embeddings.std():.4f}, mean: {embeddings.mean():.4f}")
if verbose:
# Compute the squared difference
squared_diff = (embeddings.unsqueeze(1) - embeddings.unsqueeze(0)) ** 2
squared_l2_dist = squared_diff.sum(-1)
l2_distance_matrix = torch.sqrt(squared_l2_dist)
print("Pairwise L2 Distance Matrix:")
print(" \t" + "\t".join(example_tokens))
for i, row in enumerate(l2_distance_matrix):
print(f"{example_tokens[i]}\t" + "\t".join(f"{dist:.4f}" for dist in row))
# We're working in cosine-similarity space
# So first, renormalize the embeddings to have norm 1
embedding_norms = torch.norm(embeddings, dim=-1, keepdim=True)
embeddings = embeddings / embedding_norms
print(f"embedding norms pre normalization:")
print(embedding_norms)
print(f"embedding norms post normalization:")
print(torch.norm(embeddings, dim=-1, keepdim=True))
print(f"Using {len(embeddings)} embeddings to compute initial embedding...")
init_embedding = embeddings.mean(dim=0)
# normalize the init_embedding to have norm 1:
init_embedding = init_embedding / torch.norm(init_embedding)
# rescale the init_embedding to have the same std as the average of the embeddings:
init_embedding = init_embedding * embedding_norms.mean()
print(f"init_embedding norm: {torch.norm(init_embedding):.4f}, std: {init_embedding.std():.4f}, mean: {init_embedding.mean():.4f}")
if (desired_std_multiplier is not None) and desired_std_multiplier > 0:
avg_std = torch.stack(stds).mean()
current_std = init_embedding.std()
scale_factor = desired_std_multiplier * avg_std / current_std
init_embedding = init_embedding * scale_factor
print(f"Scaled Mean Embedding: std: {init_embedding.std():.4f}, mean: {init_embedding.mean():.4f}")
return init_embedding
def plot_token_embeddings(self, example_tokens, output_folder = ".", x_range = [-0.05, 0.05]):
print(f"Plotting embeddings for tokens: {example_tokens}")
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if tokenizer is None:
idx += 1
continue
token_ids = tokenizer.convert_tokens_to_ids(example_tokens)
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_ids].clone()
# plot the embeddings histogram:
for token_name, embedding in zip(example_tokens, embeddings):
plot_torch_hist(embedding, 0, output_folder, f"tok_{token_name}_{idx}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
idx += 1
def initialize_new_tokens(self,
inserting_toks: List[str],
starting_toks: Optional[List[str]] = None,
seed: int = 0,
):
print("Initializing new tokens...")
print(inserting_toks)
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if tokenizer is None:
idx += 1
continue
assert isinstance(
inserting_toks, list
), "inserting_toks should be a list of strings."
assert all(
isinstance(tok, str) for tok in inserting_toks
), "All elements in inserting_toks should be strings."
self.inserting_toks = inserting_toks
print(f"Inserting new tokens into tokenizer-{idx}:")
print(self.inserting_toks)
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
tokenizer.add_special_tokens(special_tokens_dict)
text_encoder.resize_token_embeddings(len(tokenizer))
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
# random initialization of new tokens
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std() #(axis=1).mean()
)
std_token_mean = (
text_encoder.text_model.embeddings.token_embedding.weight.data.mean() #(axis=1).mean()
)
print(f"Text encoder {idx} token_embedding_std: {std_token_embedding}")
if starting_toks is not None:
assert isinstance(
starting_toks, list
), "starting_toks should be a list of strings."
assert all(
isinstance(tok, str) for tok in starting_toks
), "All elements in starting_toks should be strings."
assert len(starting_toks) == len(self.inserting_toks), "starting_toks should have the same length as inserting_toks"
self.starting_ids = tokenizer.convert_tokens_to_ids(starting_toks)
print(f"Copying embeddings from starting tokens {starting_toks} to new tokens {self.inserting_toks}")
print(f"Starting ids: {self.starting_ids}")
# copy the embeddings of the starting tokens to the new tokens
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
else:
if 1: # random initialization:
torch.manual_seed(seed)
init_embeddings = (torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype) * std_token_embedding * 1.0)
else:
# Test code to initialize the new tokens with some specific tokens
first_tokens = [
"Sophia",
"Liam",
"Ethan",
"Lucas",
"Olivia",
"Noah",
"John",
"David",
"James",
"Robert",
"Michael",
"William",
]
second_tokens = [
"Smith",
"Johnson",
"Williams",
"Brown",
"Jones",
"Garcia",
"Miller",
"Davis",
"Rodriguez",
"Carter",
"Trump",
"Clinton",
"Wilson",
"Harris",
"Lewis",
"Scott"
]
self.anchor_embedding_one = self.get_start_embedding(text_encoder, tokenizer, first_tokens)
self.anchor_embedding_two = self.get_start_embedding(text_encoder, tokenizer, second_tokens)
self.anchor_embedding_three = self.get_start_embedding(text_encoder, tokenizer, first_tokens)
self.anchor_embedding_four = self.get_start_embedding(text_encoder, tokenizer, second_tokens)
init_embeddings = torch.stack([self.anchor_embedding_one, self.anchor_embedding_two, self.anchor_embedding_three, self.anchor_embedding_four])
print(f"init_embedding std: {init_embeddings.std():.4f}, avg-std: {std_token_embedding:.4f}")
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
self.embeddings_settings[f"index_no_updates_{idx}"] = inu
idx += 1
def pre_optimize_token_embeddings(self, train_dataset, epochs=10):
### THIS FUNCTION IS NOT DONE YET
### Idea here is to use CLIP-similarity between imgs and prompts to pre-optimize the embeddings
for idx in range(len(train_dataset)):
(tok1, tok2), vae_latent, mask = train_dataset[idx]
image_path = train_dataset.image_path[idx]
image_path = os.path.join(os.path.dirname(train_dataset.csv_path), image_path)
image = PIL.Image.open(image_path).convert("RGB")
print(f"---> Loaded sample {idx}:")
print("Tokens:")
print(tok1.shape)
print(tok2.shape)
print("Image:")
print(image.size)
# tokens to text embeds
prompt_embeds_list = []
#for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
for tok, text_encoder in zip((tok1, tok2), self.text_encoders):
prompt_embeds_out = text_encoder(
tok.to(text_encoder.device),
output_hidden_states=True,
)
print("prompt_embeds_out:")
print(prompt_embeds_out.shape)
pooled_prompt_embeds = prompt_embeds_out[0]
prompt_embeds = prompt_embeds_out.hidden_states[-2]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
prompt_embeds_list.append(prompt_embeds)
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
print("prompt_embeds:")
print(prompt_embeds.shape)
print("pooled_prompt_embeds:")
print(pooled_prompt_embeds.shape)
def save_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
assert (
self.train_ids is not None
), "Initialize new tokens before saving embeddings."
tensors = {}
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
tensors[txt_encoder_keys[idx]] = new_token_embeddings
save_file(tensors, file_path)
@property
def dtype(self):
return self.text_encoders[0].dtype
@property
def device(self):
return self.text_encoders[0].device
def _compute_off_ratio(self, idx):
# compute the off-std-ratio for the embeddings
text_encoder = self.text_encoders[idx]
tokenizer = self.tokenizers[idx]
if text_encoder is None:
off_ratio = -1
else:
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
index_updates = ~index_no_updates
new_embeddings = (text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates])
off_ratio = std_token_embedding / new_embeddings.std()
return off_ratio
def fix_embedding_std(self, off_ratio_power = 0.1):
std_penalty = 0.0
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
index_updates = ~index_no_updates
new_embeddings = (text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates])
off_ratio = self._compute_off_ratio(idx)
std_penalty += (off_ratio - 1.0)**2
if (off_ratio < 0.95) or (off_ratio > 1.05):
print(f"std-off ratio-{idx} (target-std / embedding-std) = {off_ratio:.4f}, prob not ideal...")
print(f"std_token_embedding: {std_token_embedding}")
print(f"std new_embeddings: {new_embeddings.std()}")
# rescale the embeddings to have a more similar std as before:
new_embeddings = new_embeddings * (off_ratio**off_ratio_power)
text_encoder.text_model.embeddings.token_embedding.weight.data[
index_updates
] = new_embeddings
idx += 1
@torch.no_grad()
def retract_embeddings(self, print_stds = False):
idx = 0
means, stds = [], []
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
text_encoder.text_model.embeddings.token_embedding.weight.data[
index_no_updates
] = (
self.embeddings_settings[f"original_embeddings_{idx}"][index_no_updates]
.to(device=text_encoder.device)
.to(dtype=text_encoder.dtype)
)
# for the parts that were updated, we can normalize them a bit
# to have the same std as before
std_token_embedding = self.embeddings_settings[f"std_token_embedding_{idx}"]
index_updates = ~index_no_updates
new_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
index_updates
]
)
idx += 1
if 0:
# get the actual embeddings that will get updated:
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
updateable_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[~inu].detach().clone().to(dtype=torch.float32).cpu().numpy()
mean_0, mean_1 = updateable_embeddings[0].mean(), updateable_embeddings[1].mean()
std_0, std_1 = updateable_embeddings[0].std(), updateable_embeddings[1].std()
means.append((mean_0, mean_1))
stds.append((std_0, std_1))
if print_stds:
print(f"Text Encoder {idx} token embeddings:")
print(f" --- Means: ({mean_0:.6f}, {mean_1:.6f})")
print(f" --- Stds: ({std_0:.6f}, {std_1:.6f})")
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
# Assuming new tokens are of the format <s_i>
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
tokenizer.add_special_tokens(special_tokens_dict)
text_encoder.resize_token_embeddings(len(tokenizer))
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
assert self.train_ids is not None, "New tokens could not be converted to IDs."
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
if not os.path.exists(file_path):
file_path = file_path.replace(".pti", ".safetensors")
if not os.path.exists(file_path):
raise FileNotFoundError(f"{file_path} does not exist.")
with safe_open(file_path, framework="pt", device=self.device.type) as f:
for idx in range(len(self.text_encoders)):
text_encoder = self.text_encoders[idx]
tokenizer = self.tokenizers[idx]
if text_encoder is None:
continue
try:
loaded_embeddings = f.get_tensor(txt_encoder_keys[idx])
except:
loaded_embeddings = f.get_tensor(f"text_encoders_{idx}")
self._load_embeddings(loaded_embeddings, tokenizer, text_encoder)
+497
View File
@@ -0,0 +1,497 @@
import os
import torch
import torch.nn.functional as F
import numpy as np
import PIL
from tqdm import tqdm
from typing import List, Optional, Dict
from safetensors.torch import save_file, safe_open
import matplotlib.pyplot as plt
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
class TokenEmbeddingsHandler:
def __init__(self, text_encoders, tokenizers):
self.text_encoders = text_encoders
self.tokenizers = tokenizers
self.train_ids: Optional[torch.Tensor] = None
self.inserting_toks: Optional[List[str]] = None
self.embeddings_settings = {}
self.target_prompt = ""
self.token_regularizer = None
def make_embeddings_trainable(self):
"""
Sets requires_grad to True for specific indices directly in the embeddings weight tensor.
"""
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
# Directly accessing and modifying the original weights tensor
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
def get_trainable_embeddings(self):
return self.get_embeddings_and_tokens(self.train_ids)
def get_embeddings_and_tokens(self, indices):
"""
Get the embeddings and tokens for the given indices using PyTorch indexing.
This version avoids detaching the original tensor and returns a view into the original
weights tensor whenever possible.
"""
embeddings, tokens = {}, {}
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
embeddings[f'txt_encoder_{idx}'] = token_embeddings
# Get all corresponding tokens for these embeddings
token_list = self.tokenizers[idx].convert_ids_to_tokens(indices)
tokens[f'txt_encoder_{idx}'] = token_list
return embeddings, tokens
def visualize_random_token_embeddings(self, output_dir, n = 6, token_list = None):
"""
Visualize the embeddings of n random tokens from each text encoder
"""
if token_list is not None:
# Convert tokens to indices using the first tokenizer
indices = self.tokenizers[0].convert_tokens_to_ids(token_list)
else:
# Randomly select indices
n_tokens = len(self.text_encoders[0].text_model.embeddings.token_embedding.weight.data)
indices = np.random.randint(0, n_tokens, n)
embeddings, tokens = self.get_embeddings_and_tokens(indices)
# Visualize the embeddings:
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
for i in range(n):
token = tokens[f'txt_encoder_{idx}'][i]
# Strip any backslashes from the token name:
token = token.replace("/", "_")
embedding = embeddings[f'txt_encoder_{idx}'][i]
plot_torch_hist(embedding, 0, os.path.join(output_dir, 'ti_embeddings') , f"frozen_enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'green')
def find_nearest_tokens(self, query_embedding, tokenizer, text_encoder, idx, distance_metric, top_k = 5):
# given a query embedding, compute the distance to all embeddings in the text encoder
# and return the top_k closest tokens
assert distance_metric in ["l2", "cosine"], "distance_metric should be either 'l2' or 'cosine'"
# get all non-optimized embeddings:
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[index_no_updates]
# compute the distance between the query embedding and all embeddings:
if distance_metric == "l2":
diff = (embeddings - query_embedding.unsqueeze(0))**2
distances = diff.sum(-1)
distances, indices = torch.topk(distances, top_k, dim=0, largest=False)
elif distance_metric == "cosine":
distances = F.cosine_similarity(embeddings, query_embedding.unsqueeze(0), dim=-1)
distances, indices = torch.topk(distances, top_k, dim=0, largest=True)
nearest_tokens = tokenizer.convert_ids_to_tokens(indices)
return nearest_tokens, distances
def print_token_info(self, distance_metric = "cosine"):
print(f"----------- Closest tokens (distance_metric = {distance_metric}) --------------")
current_token_embeddings, current_tokens = self.get_trainable_embeddings()
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
query_embeddings = current_token_embeddings[f'txt_encoder_{idx}']
query_tokens = current_tokens[f'txt_encoder_{idx}']
for token_id, query_embedding in enumerate(query_embeddings):
nearest_tokens, distances = self.find_nearest_tokens(query_embedding, tokenizer, text_encoder, idx, distance_metric)
# print the results:
print(f"txt-encoder {idx}, token {token_id}: {query_tokens[token_id]}:")
for i, (token, dist) in enumerate(zip(nearest_tokens, distances)):
print(f"---> {distance_metric} of {dist:.4f}: {token}")
idx += 1
def plot_token_embeddings(self, example_tokens, output_folder = ".", x_range = [-0.05, 0.05]):
print(f"Plotting embeddings for tokens: {example_tokens}")
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if tokenizer is None:
idx += 1
continue
token_ids = tokenizer.convert_tokens_to_ids(example_tokens)
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_ids].clone()
# plot the embeddings histogram:
for token_name, embedding in zip(example_tokens, embeddings):
plot_torch_hist(embedding, 0, output_folder, f"tok_{token_name}_{idx}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
idx += 1
@property
def dtype(self):
return self.text_encoders[0].dtype
def initialize_new_tokens(self,
inserting_toks: List[str],
starting_toks: Optional[List[str]] = None,
seed: int = 0,
):
assert isinstance(
inserting_toks, list
), "inserting_toks should be a list of strings."
assert all(
isinstance(tok, str) for tok in inserting_toks
), "All elements in inserting_toks should be strings."
print(f"Initializing new tokens: {inserting_toks}")
self.inserting_toks = inserting_toks
seed_everything(seed)
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if tokenizer is None:
idx += 1
continue
print(f"Inserting new tokens into tokenizer-{idx}:")
print(self.inserting_toks)
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
tokenizer.add_special_tokens(special_tokens_dict)
text_encoder.resize_token_embeddings(len(tokenizer))
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
# construct the indices for all the non-trainable embeddings:
all_indices = torch.linspace(0, len(tokenizer) - 1, len(tokenizer), dtype=torch.long)
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
self.non_train_ids = all_indices[inu]
# random initialization of new tokens
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
if starting_toks is not None:
assert len(starting_toks) == len(self.inserting_toks), "starting_toks should have the same length as inserting_toks"
self.starting_ids = tokenizer.convert_tokens_to_ids(starting_toks)
print(f"Copying embeddings from starting tokens {starting_toks} to new tokens {self.inserting_toks}")
print(f"Starting ids: {self.starting_ids}")
# copy the embeddings of the starting tokens to the new tokens
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
else:
std_multiplier = 1.0
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
current_std = init_embeddings.std(dim=1).mean()
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
self.embeddings_settings[f"index_no_updates_{idx}"] = inu
idx += 1
def plot_tokenid(self, token_id, suffix = '', output_folder = ".", x_range = [-0.05, 0.05]):
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if tokenizer is None:
idx += 1
continue
embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data[token_id].clone()
plot_torch_hist(embeddings, 0, output_folder, f"tok_{token_id}_{idx}_{suffix}", bins=100, min_val=x_range[0], max_val=x_range[1], ymax_f = 0.05)
idx += 1
def get_conditioning_signals(self, config, pipe, captions):
conditioning_signals = pipe.encode_prompt(
prompt=captions,
device=pipe.unet.device,
num_images_per_prompt=1,
do_classifier_free_guidance=True,
negative_prompt=None,
clip_skip=None,
)
try: # sd15
prompt_embeds, negative_prompt_embeds = conditioning_signals
pooled_prompt_embeds, add_time_ids = None, None
except: # sdxl
(
prompt_embeds,
negative_prompt_embeds,
pooled_prompt_embeds,
negative_pooled_prompt_embeds,
) = conditioning_signals
# Create Spatial-dimensional conditions.
# I dont understand why, but I get better results hardcoding the original_size values...
# original_size = (config.resolution, config.resolution)
original_size = (1024, 1024)
target_size = (config.resolution, config.resolution)
crops_coords_top_left = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
else:
text_encoder_projection_dim = pipe.text_encoder_2.config.projection_dim
add_time_ids = pipe._get_add_time_ids(
original_size,
crops_coords_top_left,
target_size,
dtype=prompt_embeds.dtype,
text_encoder_projection_dim=text_encoder_projection_dim,
)
add_time_ids = add_time_ids.to(config.device, dtype=prompt_embeds.dtype).repeat(
prompt_embeds.shape[0], 1
)
return prompt_embeds, pooled_prompt_embeds, add_time_ids
def encode_text(self, text, config, pipe):
prompt_embeds, pooled_prompt_embeds, add_time_ids = self.get_conditioning_signals(config, pipe, [text])
return prompt_embeds, pooled_prompt_embeds
def compute_target_prompt_loss(self, target_prompt, prompt_embeds, pooled_prompt_embeds, config, pipe):
"""
Compute a distance loss between the prompt embeddings and the target prompt embeddings
"""
if target_prompt != self.target_prompt:
self.target_prompt = target_prompt
self.target_prompt_embeds, self.target_pooled_prompt_embeds = self.encode_text(self.target_prompt, config, pipe)
# detach the target prompt embeddings (we don't need gradients here, these are just static targets)
self.target_prompt_embeds = self.target_prompt_embeds.detach()
try:
self.target_pooled_prompt_embeds = self.target_pooled_prompt_embeds.detach()
except:
self.target_pooled_prompt_embeds = None
# compute the losses:
# Replicate target embeddings to match the batch size of prompt_embeds
batch_size = prompt_embeds.size(0)
target = self.target_prompt_embeds.expand(batch_size, -1, -1)
embeds_l2_loss = F.mse_loss(prompt_embeds, target)
embeds_cosine_loss = 1.0 - F.cosine_similarity(prompt_embeds, target, dim=-1).mean()
loss = embeds_l2_loss + embeds_cosine_loss
if pooled_prompt_embeds is not None:
target = self.target_pooled_prompt_embeds.expand(batch_size, -1)
pooled_embeds_l2_loss = F.mse_loss(pooled_prompt_embeds, target)
pooled_embeds_cosine_loss = 1.0 - F.cosine_similarity(pooled_prompt_embeds, target, dim=-1).mean()
loss += 0.25 * (pooled_embeds_l2_loss + pooled_embeds_cosine_loss)
return loss
def pre_optimize_token_embeddings(self, config, pipe):
"""
Warmup the token embeddings by optimizing them without using the image denoiser,
but simply using CLIP-txt and CLIP-img similarity losses
TODO: add CLIP-img similarity loss into this mix
--> This requires loading the img-encoder part for each of the txt-encoders and figuring out the correct projection layer
"""
target_prompt = config.training_attributes["gpt_description"]
if config.token_warmup_steps <= 0 or not target_prompt:
print("Skipping token embedding warmup.")
return
print(f'Warming up token embeddings with prompt: {target_prompt}...')
# Setup the token optimizer:
ti_parameters = []
for text_encoder in self.text_encoders:
if text_encoder is not None:
text_encoder.train()
for name, param in text_encoder.named_parameters():
if "token_embedding" in name:
param.requires_grad = True
ti_parameters.append(param)
params_to_optimize_ti = [{
"params": ti_parameters,
"lr": config.ti_lr,
"weight_decay": config.ti_weight_decay,
}]
optimizer_ti = torch.optim.AdamW(
params_to_optimize_ti,
weight_decay=config.ti_weight_decay,
)
token_string = config.token_dict["TOK"]
# TODO: check if some light prompt template augmentation is useful here to make the optimization more robust
prompt_template = [
'{}',
'{}',
'{}',
#'a {}',
#'{} image',
#'a picture of {}',
]
losses = {'concept_description_loss': [], 'covariance_tok_reg_loss': [], 'token_std_loss': []}
for step in tqdm(range(config.token_warmup_steps)):
if step % 30 == 0 and config.debug and 0: # disalbe this for now
for i, token_index in enumerate(self.train_ids):
self.plot_tokenid(token_index, suffix = f'token_{i}_{step}', output_folder = f'{config.output_dir}/token_opt')
# pick a random prompt template and inject the token string:
prompt_to_optimize = np.random.choice(prompt_template).format(token_string)
prompt_embeds, pooled_prompt_embeds = self.encode_text(prompt_to_optimize, config, pipe)
# Compute the target_prompt distance loss:
loss = 0.2 * self.compute_target_prompt_loss(target_prompt, prompt_embeds, pooled_prompt_embeds, config, pipe)
losses['concept_description_loss'].append(loss.item())
# Compute token regularization loss:
loss, losses, _ = self.token_regularizer.apply_regularization(loss, losses, None, prompt_embeds, std_loss_w = 0.5)
# Backward pass:
retain_graph = step < (config.token_warmup_steps - 1) # Retain graph for all but the last step
loss.backward(retain_graph=retain_graph)
# zero out the gradients of the non-trained text-encoder embeddings
for embedding_tensor in ti_parameters:
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
optimizer_ti.step()
self.fix_embedding_std(config.off_ratio_power)
optimizer_ti.zero_grad()
if config.debug:
plot_loss(losses, save_path=f'{config.output_dir}/token_warmup_loss.png')
def save_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
assert (
self.train_ids is not None
), "Initialize new tokens before saving embeddings."
# Create a set of indices for the non-train_ids:
self.not_train_ids = torch.linspace(0, len(self.tokenizers[0]) - 1, len(self.tokenizers[0]), dtype=torch.long)
tensors = {}
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
tensors[txt_encoder_keys[idx]] = new_token_embeddings
save_file(tensors, file_path)
@property
def device(self):
return self.text_encoders[0].device
def fix_embedding_std(self, off_ratio_power=0.1):
if off_ratio_power == 0.0:
return
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
# Get the standard deviation target and current embeddings.
target_std = self.embeddings_settings[f"std_token_embedding_{idx}"]
embeddings, _ = self.get_trainable_embeddings()
new_embeddings = embeddings[f'txt_encoder_{idx}']
assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!"
new_stds = new_embeddings.std(dim=1)
#off_ratios = target_std.float() / new_stds.float()
off_ratios = target_std / new_stds
# Check if off_ratios are within an acceptable range.
if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1):
# Convert the pytorch tensor into a list of python floats:
off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3)
print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...")
# Adjust embeddings using the computed ratios.
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
index_updates = ~index_no_updates
multiplier_values = off_ratios**off_ratio_power
multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings)
text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values
idx += 1
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
# Assuming new tokens are of the format <s_i>
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
special_tokens_dict = {"additional_special_tokens": self.inserting_toks}
tokenizer.add_special_tokens(special_tokens_dict)
text_encoder.resize_token_embeddings(len(tokenizer))
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
assert self.train_ids is not None, "New tokens could not be converted to IDs."
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
if not os.path.exists(file_path):
file_path = file_path.replace(".pti", ".safetensors")
if not os.path.exists(file_path):
raise FileNotFoundError(f"{file_path} does not exist.")
with safe_open(file_path, framework="pt", device=self.device.type) as f:
for idx in range(len(self.text_encoders)):
text_encoder = self.text_encoders[idx]
tokenizer = self.tokenizers[idx]
if text_encoder is None:
continue
try:
loaded_embeddings = f.get_tensor(txt_encoder_keys[idx])
except:
loaded_embeddings = f.get_tensor(f"text_encoders_{idx}")
self._load_embeddings(loaded_embeddings, tokenizer, text_encoder)
+489
View File
@@ -0,0 +1,489 @@
import torch
import os
import random
import shutil
import json
import gc
import re
from diffusers import EulerDiscreteScheduler
from trainer.utils.val_prompts import val_prompts
from trainer.utils.utils import fix_prompt, replace_in_string
from trainer.models import load_models
from .checkpoint import load_checkpoint, set_adapter_scales
from diffusers import (
DDPMScheduler,
EulerDiscreteScheduler,
StableDiffusionPipeline,
StableDiffusionXLPipeline,
)
def load_model(pretrained_model: dict):
if pretrained_model["version"] == "sd15":
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model["path"], torch_dtype=torch.float16, use_safetensors=True
)
else:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model["path"], torch_dtype=torch.float16, use_safetensors=True
)
pipe = pipe.to("cuda", dtype=torch.float16)
pipe.scheduler = EulerDiscreteScheduler.from_config(
pipe.scheduler.config
) # , timestep_spacing="trailing")
return pipe
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
"""
This function is rather ugly, but implements a custom token-replacement policy we adopted at Eden:
Basically you trigger the lora with a token "TOK" or "<concept>", and then this token gets replaced with the actual learned tokens
"""
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError(
"This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!"
)
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
trigger_text = training_args["training_attributes"]["trigger_text"]
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
lora_name = "concept"
lora_name_encapsulated = "<" + lora_name + ">"
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK",
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(
prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"}
)
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
prompt = fix_prompt(prompt)
if verbose:
print("-------------------------")
print("Adjusted prompt for LORA:")
print(orig_prompt)
print("-- to:")
print(prompt)
print("-------------------------")
return prompt
def get_conditioning_signals(config, pipe, captions):
conditioning_signals = pipe.encode_prompt(
prompt=captions,
device=pipe.unet.device,
num_images_per_prompt=1,
do_classifier_free_guidance=True,
negative_prompt=None,
clip_skip=None,
)
try: # sd15
prompt_embeds, negative_prompt_embeds = conditioning_signals
pooled_prompt_embeds, add_time_ids = None, None
except: # sdxl
(
prompt_embeds,
negative_prompt_embeds,
pooled_prompt_embeds,
negative_pooled_prompt_embeds,
) = conditioning_signals
# Create Spatial-dimensional conditions.
# I dont understand why, but I get better results hardcoding the original_size values...
# original_size = (config.resolution, config.resolution)
original_size = (1024, 1024)
target_size = (config.resolution, config.resolution)
crops_coords_top_left = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
else:
text_encoder_projection_dim = pipe.text_encoder_2.config.projection_dim
add_time_ids = pipe._get_add_time_ids(
original_size,
crops_coords_top_left,
target_size,
dtype=prompt_embeds.dtype,
text_encoder_projection_dim=text_encoder_projection_dim,
)
add_time_ids = add_time_ids.to(config.device, dtype=prompt_embeds.dtype).repeat(
prompt_embeds.shape[0], 1
)
return prompt_embeds, pooled_prompt_embeds, add_time_ids
def blend_conditions(
embeds1,
embeds2,
lora_scale,
token_scale_power=0.4, # adjusts the curve of the interpolation
min_token_scale=0.5, # minimum token scale (corresponds to lora_scale = 0)
token_scale=None,
verbose=1,
):
"""
using lora_scale, apply linear interpolation between two sets of embeddings
"""
try: # sdxl:
c1, uc1, pc1, puc1 = embeds1
c2, uc2, pc2, puc2 = embeds2
except: # sd15:
c1, uc1 = embeds1
c2, uc2 = embeds2
pc1, pc2, puc1, puc2 = None, None, None, None
if token_scale is None: # compute the token_scale based on lora_scale:
token_scale = lora_scale**token_scale_power
# rescale the [0,1] range to [min_token_scale, 1] range:
token_scale = min_token_scale + (1 - min_token_scale) * token_scale
if verbose:
print(
f"Setting token_scale to {token_scale:.2f} (lora_scale = {lora_scale:.2f}, power = {token_scale_power})"
)
try:
c = (1 - token_scale) * c1 + token_scale * c2
uc = (1 - token_scale) * uc1 + token_scale * uc2
try:
pc = (1 - token_scale) * pc1 + token_scale * pc2
puc = (1 - token_scale) * puc1 + token_scale * puc2
except:
pc, puc = None, None
embeds = (c, uc, pc, puc)
except:
print(
f"Error in blending conditions for toking interpolation, falling back to embeds2"
)
token_scale = 1.0
embeds = (c2, uc2, pc2, puc2)
return embeds, token_scale
def encode_prompt_advanced(
pipe,
lora_path,
prompt,
negative_prompt,
lora_scale,
guidance_scale,
token_scale=None,
concept_mode=None,
):
"""
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
This allows interpolating the strength of the trained token in the final image.
"""
if lora_path:
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
else:
lora_prompt = prompt
if concept_mode == "face":
replace_str = "person"
elif concept_mode == "object":
replace_str = "object"
else:
replace_str = ""
zero_prompt = prompt.replace("<concept>", replace_str)
zero_prompt = fix_prompt(zero_prompt)
print(f"Embedding lora prompt: {lora_prompt}")
print(f"Embedding zero prompt: {zero_prompt}")
try: # sdxl:
embeds = pipe.encode_prompt(
lora_prompt,
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt,
)
zero_embeds = pipe.encode_prompt(
zero_prompt,
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt,
)
except: # sd15:
embeds = pipe.encode_prompt(lora_prompt, pipe.device, 1, True, negative_prompt)
zero_embeds = pipe.encode_prompt(
zero_prompt, pipe.device, 1, True, negative_prompt
)
embeds, token_scale = blend_conditions(
zero_embeds, embeds, lora_scale, token_scale=token_scale
)
return embeds
@torch.no_grad()
def render_images(
render_size,
lora_path,
train_step,
seed,
is_lora,
pretrained_model,
lora_scale,
n_steps=25,
n_imgs=4,
device="cuda:0",
pipe = None,
checkpoint_folder: str = None,
):
if checkpoint_folder is not None:
assert pipe is None, f"Expected either one of checkpoint_folder or pipe to be None. But got: checkpoint_folder: {checkpoint_folder} and pipe is not None"
if pipe is not None:
assert checkpoint_folder is None, f"Expected either one of checkpoint_folder or pipe to be None. But got pipe is NOT None checkpoint_folder is: {checkpoint_folder}"
random.seed(seed)
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
concept_mode = training_args["concept_mode"]
if concept_mode == "style":
validation_prompts_raw = random.sample(val_prompts["style"], n_imgs)
validation_prompts_raw[0] = ""
elif concept_mode == "face":
validation_prompts_raw = random.sample(val_prompts["face"], n_imgs)
validation_prompts_raw[0] = "<concept>"
else:
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
validation_prompts_raw[0] = "<concept>"
if (
checkpoint_folder is not None
): # reload the entire pipeline from disk and load in the lora module
print(f"Reloading checkpoint from disk: {checkpoint_folder}")
gc.collect()
torch.cuda.empty_cache()
pipe = load_checkpoint(
pretrained_model_version=pretrained_model["version"],
pretrained_model_path=pretrained_model["path"],
lora_save_path=checkpoint_folder,
is_lora=is_lora,
device=device,
lora_scale=lora_scale,
)
else:
assert pipe is not None
training_scheduler = pipe.scheduler
print(f"Using existing model for inference")
print(
f"Re-using training pipeline for inference, just swapping the scheduler.."
)
pipe.vae = pipe.vae.to(device).to(pipe.unet.dtype)
pipe = set_adapter_scales(pipe, lora_scale = lora_scale)
pipe.scheduler = EulerDiscreteScheduler.from_config(
pipe.scheduler.config, timestep_spacing="trailing"
)
generator = torch.Generator(device=device).manual_seed(seed)
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
pipeline_args = {
"num_inference_steps": n_steps,
"guidance_scale": 8,
"width": render_size[0],
"height": render_size[1],
}
for i in range(n_imgs):
print(f"Rendering validation img with prompt: {validation_prompts_raw[i]}")
c, uc, pc, puc = encode_prompt_advanced(
pipe,
lora_path,
validation_prompts_raw[i],
negative_prompt,
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
)
pipeline_args["prompt_embeds"] = c
pipeline_args["negative_prompt_embeds"] = uc
if pretrained_model["version"] == "sdxl":
pipeline_args["pooled_prompt_embeds"] = pc
pipeline_args["negative_pooled_prompt_embeds"] = puc
image = pipe(**pipeline_args, generator=generator).images[0]
image.save(
os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"),
format="JPEG",
quality=95,
)
if checkpoint_folder is None:
pipe.scheduler = training_scheduler
pipe.vae = pipe.vae.to("cpu")
gc.collect()
torch.cuda.empty_cache()
# reset the adapter scales to 1.0
pipe = set_adapter_scales(pipe, lora_scale = 1.0)
return validation_prompts_raw
@torch.no_grad()
def render_images_eval(
concept_mode: str,
output_folder: str,
render_size: tuple,
checkpoint_folder: str,
seed: int,
is_lora: bool,
pretrained_model: dict,
trigger_text: str,
lora_scale=0.7,
n_steps=25,
n_imgs=4,
device="cuda:0",
verbose: bool = True,
):
random.seed(seed)
assert os.path.exists(output_folder), f"Invalid folder: {output_folder}"
if concept_mode == "style":
validation_prompts_raw = random.sample(val_prompts["style"], n_imgs)
validation_prompts_raw[0] = ""
elif concept_mode == "face":
validation_prompts_raw = random.sample(val_prompts["face"], n_imgs)
validation_prompts_raw[0] = "<concept>"
else:
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
validation_prompts_raw[0] = "<concept>"
print(f"Reloading entire pipeline from disk for eval...")
gc.collect()
torch.cuda.empty_cache()
pipe = load_checkpoint(
pretrained_model_version=pretrained_model["version"],
pretrained_model_path=pretrained_model["path"],
lora_save_path=checkpoint_folder,
is_lora=is_lora,
device=device,
)
pipe.scheduler = EulerDiscreteScheduler.from_config(
pipe.scheduler.config, timestep_spacing="trailing"
)
generator = torch.Generator(device=device).manual_seed(seed)
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
pipeline_args = {
"num_inference_steps": n_steps,
"guidance_scale": 8,
"height": render_size[0],
"width": render_size[1],
}
filenames = []
for i in range(n_imgs):
print(f"Rendering validation img with prompt: {validation_prompts_raw[i]}")
c, uc, pc, puc = encode_prompt_advanced(
pipe,
checkpoint_folder,
validation_prompts_raw[i],
negative_prompt,
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
)
pipeline_args["prompt_embeds"] = c
pipeline_args["negative_prompt_embeds"] = uc
if pretrained_model["version"] == "sdxl":
pipeline_args["pooled_prompt_embeds"] = pc
pipeline_args["negative_pooled_prompt_embeds"] = puc
image = pipe(**pipeline_args, generator=generator).images[0]
filename = os.path.join(output_folder, f"{i}.jpg")
image.save(
filename,
format="JPEG",
quality=95,
)
filenames.append(filename)
return filenames, validation_prompts_raw
+375
View File
@@ -0,0 +1,375 @@
import os
import time
import matplotlib.pyplot as plt
import torch
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
from trainer.inference import get_conditioning_signals
def compute_snr(noise_scheduler, timesteps):
"""
Computes SNR as per
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
"""
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = alphas_cumprod**0.5
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
# Expand the tensors.
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
# Compute SNR.
snr = (alpha / sigma) ** 2
return snr
def compute_grad_norm(parameters, norm_type = 2.0, foreach = None, error_if_nonfinite = False):
if isinstance(parameters, torch.Tensor):
parameters = [parameters]
grads = [p.grad for p in parameters if p.grad is not None]
first_device = grads[0].device
grouped_grads = _group_tensors_by_device_and_dtype([[g.detach() for g in grads]])
norms = []
for ((device, _), ([grads], _)) in grouped_grads.items():
if (foreach is None or foreach) and _has_foreach_support(grads, device=device):
norms.extend(torch._foreach_norm(grads, norm_type))
elif foreach:
raise RuntimeError(f'foreach=True was passed, but can\'t use the foreach API on {device.type} tensors')
else:
norms.extend([torch.linalg.vector_norm(g, norm_type) for g in grads])
total_norm = torch.linalg.vector_norm(torch.stack([norm.to(first_device) for norm in norms]), norm_type)
return total_norm
def compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps):
# Get the unet prediction target depending on the prediction type:
if noise_scheduler.config.prediction_type == "epsilon":
target = noise
elif noise_scheduler.config.prediction_type == "v_prediction":
print(f"Using velocity prediction!")
target = noise_scheduler.get_velocity(noisy_latent, noise, timesteps)
else:
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
loss = (model_pred - target).pow(2) * mask
if config.snr_gamma is None or config.snr_gamma == 0.0:
# modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
loss = loss.mean()
else:
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
# This is discussed in Section 4.2 of the same paper.
snr = compute_snr(noise_scheduler, timesteps)
base_weight = (
torch.stack([snr, config.snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
)
if noise_scheduler.config.prediction_type == "v_prediction":
# Velocity objective needs to be floored to an SNR weight of one.
mse_loss_weights = base_weight + 1
else:
# Epsilon and sample both use the same loss weights.
mse_loss_weights = base_weight
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
# modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
loss = loss.mean()
return loss
class ConditioningRegularizer:
"""
Regularizes:
- the norms of the prompt_conditioning vectors
- the statistics of the token embeddings.
"""
def __init__(self, config, embedding_handler):
self.config = config
self.embedding_handler = embedding_handler
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
self.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
self.distribution_regularizers = {}
idx = 0
for tokenizer, text_encoder in zip(embedding_handler.tokenizers, embedding_handler.text_encoders):
if tokenizer is None:
idx += 1
continue
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
idx += 1
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.003, pipe=None):
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
if self.config.cond_reg_w > 0.0:
reg_loss, regularization_norm_value = self._compute_regularization_loss(prompt_embeds)
loss += self.config.cond_reg_w * reg_loss
if prompt_embeds_norms is not None:
prompt_embeds_norms['main'].append(regularization_norm_value.item())
if self.config.tok_cond_reg_w > 0.0 and pipe is not None:
reg_loss, regularization_norm_value = self._compute_tok_regularization_loss(pipe)
loss += self.config.tok_cond_reg_w * reg_loss
if prompt_embeds_norms is not None:
prompt_embeds_norms['reg'].append(regularization_norm_value.item())
if self.config.tok_cov_reg_w > 0.0:
tot_reg_losses = []
for key, distribution_regularizer in self.distribution_regularizers.items():
reg_loss = distribution_regularizer.compute_covariance_loss(self.embedding_handler.get_trainable_embeddings()[0][key])
tot_reg_losses.append(reg_loss)
mean_reg_loss = torch.stack(tot_reg_losses).mean()
loss += self.config.tok_cov_reg_w * mean_reg_loss
losses['covariance_tok_reg_loss'].append(mean_reg_loss.item())
if std_loss_w > 0.0:
tot_std_losses = []
for key, distribution_regularizer in self.distribution_regularizers.items():
std_loss = distribution_regularizer.compute_std_loss(self.embedding_handler.get_trainable_embeddings()[0][key])
tot_std_losses.append(std_loss)
mean_std_loss = torch.stack(tot_std_losses).mean()
loss += std_loss_w * mean_std_loss
losses['token_std_loss'].append(mean_std_loss.item())
return loss, losses, prompt_embeds_norms
def _compute_regularization_loss(self, prompt_embeds):
conditioning_norms = prompt_embeds.norm(dim=-1).mean(dim=0)
regularization_norm_value = conditioning_norms[2:].mean()
reg_loss = (regularization_norm_value - self.target_norm).pow(2)
return reg_loss, regularization_norm_value
def _compute_tok_regularization_loss(self, pipe):
reg_captions = [caption.replace("TOK", self.token_replacement) for caption in self.reg_captions]
reg_prompt_embeds, reg_pooled_prompt_embeds, reg_add_time_ids = get_conditioning_signals(
self.config, pipe, reg_captions
)
reg_conditioning_norms = reg_prompt_embeds.norm(dim=-1).mean(dim=0)
regularization_norm_value = reg_conditioning_norms[2:].mean()
reg_loss = (regularization_norm_value - self.target_norm).pow(2)
return reg_loss, regularization_norm_value
class DistributionLoss(torch.nn.Module):
"""
Class to simplify the calculation of the covariance loss between the trained token embeddings and the pretrained embeddings.
"""
def __init__(self, pretrained_embeddings, dtype=torch.float32, outdir = None):
super(DistributionLoss, self).__init__()
print(f"Initialized a new DistributionLoss with shape: {pretrained_embeddings.shape}")
self.dtype = dtype
self.target_cov = self._calculate_covariance(pretrained_embeddings)
self.target_stds = pretrained_embeddings.std(-1)
self.target_stds_mean = self.target_stds.mean()
self.target_stds_var = self.target_stds.std()**2 / self.target_stds.mean()
if outdir:
# Plot a histogram of the stds:
plt.figure()
plt.hist(self.target_stds.detach().float().cpu().numpy(), bins=100)
plt.title(f"stds of tokens (shape = {pretrained_embeddings.shape[0]} x {pretrained_embeddings.shape[1]})")
plt.xlim(0, 0.02)
plt.savefig(os.path.join(outdir, f"stds_histogram_{int(time.time()*100)}.png"))
def _calculate_covariance(self, embeddings):
embeddings = embeddings.to(self.dtype)
mean = embeddings.mean(0)
embeddings_adjusted = embeddings - mean
covariance = torch.mm(embeddings_adjusted.T, embeddings_adjusted) / (embeddings.size(0) - 1)
return covariance
def compute_covariance_loss(self, new_embeddings):
input_dtype = new_embeddings.dtype
cov_new = self._calculate_covariance(new_embeddings)
# Normalizing by the product of the dimensions of the covariance matrix.
num_features = new_embeddings.size(1) # Assuming embeddings are of shape [n_samples, n_features]
scale_factor = num_features * num_features
loss = torch.norm(self.target_cov - cov_new, p='fro') / scale_factor
return loss.to(input_dtype)
def compute_std_loss(self, new_embeddings):
if new_embeddings.size(1) == 1:
new_embeddings = new_embeddings.unsqueeze(0)
deviation_loss = ((self.target_stds_mean - new_embeddings.std(-1))**2 / self.target_stds_var).mean()
return deviation_loss
#######################################################################
#######################################################################
## Everything below here is experimental stuff not yet fully functional:
import torch
import numpy as np
from torch.distributions import MultivariateNormal, Normal
from torch.distributions.distribution import Distribution
class GaussianKDE(Distribution):
def __init__(self, X, bw = 0.1):
"""
X : tensor (n, d)
`n` points with `d` dimensions to which KDE will be fit
bw : numeric
bandwidth for Gaussian kernel
"""
self.X = X
self.bw = bw
self.dims = X.shape[-1]
self.n = X.shape[0]
self.mvn = MultivariateNormal(loc=torch.zeros(self.dims),
covariance_matrix=torch.eye(self.dims))
def sample(self, num_samples):
idxs = (np.random.uniform(0, 1, num_samples) * self.n).astype(int)
norm = Normal(loc=self.X[idxs], scale=self.bw)
return norm.sample()
def score_samples(self, Y, X=None):
"""Returns the kernel density estimates of each point in `Y`.
Parameters
----------
Y : tensor (m, d)
`m` points with `d` dimensions for which the probability density will
be calculated
X : tensor (n, d), optional
`n` points with `d` dimensions to which KDE will be fit. Provided to
allow batch calculations in `log_prob`. By default, `X` is None and
all points used to initialize KernelDensityEstimator are included.
Returns
-------
log_probs : tensor (m)
log probability densities for each of the queried points in `Y`
"""
if X == None:
X = self.X
log_probs = torch.log(
(self.bw**(-self.dims) *
torch.exp(self.mvn.log_prob(
(X.unsqueeze(1) - Y) / self.bw))).sum(dim=0) / self.n)
return log_probs
def log_prob(self, Y):
"""Returns the total log probability of one or more points, `Y`, using
a Multivariate Normal kernel fit to `X` and scaled using `bw`.
Parameters
----------
Y : tensor (m, d)
`m` points with `d` dimensions for which the probability density will
be calculated
Returns
-------
log_prob : numeric
total log probability density for the queried points, `Y`
"""
X_chunks = self.X.split(1000)
Y_chunks = Y.split(1000)
log_prob = 0
for x in X_chunks:
for y in Y_chunks:
log_prob += self.score_samples(y, x).sum(dim=0)
return log_prob
class DifferentiableHistogram:
"""
TODO fix this function
"""
def __init__(self, x, bins=64, min_range=None, max_range=None, bandwidth=0.02):
self.bins = bins
self.bandwidth = bandwidth * (x.max() - x.min())
if min_range is None or max_range is None:
self.min_range = x.min()
self.max_range = x.max()
else:
self.min_range = min_range
self.max_range = max_range
# Create bins
self.bin_edges = torch.linspace(self.min_range, self.max_range, bins + 1).to(x.device)
self.bin_centers = (self.bin_edges[:-1] + self.bin_edges[1:]) / 2.0
# Compute histogram using Gaussian smoothing with bandwidth
distances = (x.unsqueeze(1) - self.bin_centers.unsqueeze(0)) / self.bandwidth
weights = torch.exp(-0.5 * (distances ** 2))
histogram = weights.sum(dim=0)
# Normalize to form a PDF
self.pdf = histogram / histogram.sum()
# Plot the histogram for validation
plt.figure()
plt.plot(self.bin_centers.float().cpu().numpy(), self.pdf.float().cpu().numpy())
plt.title(f"PDF of token embeddings (shape = {x.shape})")
plt.xlim(0, x.max().item()*1.1)
plt.savefig(f"pdf_histogram_{int(time.time()*100)}.png")
plt.close()
def __call__(self, y):
"""
Compute the negative log likelihood for a given sample y.
Arguments:
- y: Tensor of shape (m,) for which to compute the loss.
Returns:
- loss: Scalar representing the negative log likelihood of sample y.
"""
y_distances = (y.unsqueeze(1) - self.bin_centers.unsqueeze(0)) / self.bandwidth
y_weights = torch.exp(-0.5 * (y_distances ** 2))
likelihoods = (self.pdf * y_weights).sum(dim=1)
nll = -torch.log(likelihoods).mean()
return nll
"""
Learned Notes on the token embeddings:
shape = 49410, 768
Computed means of [1, 768] = [0, 0, 0, 0,...]
Computed stds of [1, 768] = [0.0139, 0.0139, 0.0139, ...]
Computed means of [49410, 1] = [0, 0, 0, 0,...]
Computed stds of [49410, 1] = [0.0151, 0.0154, 0.0141, ..., 0.0396, 0.0150, 0.0148]
"""
+97
View File
@@ -0,0 +1,97 @@
import os
import time
import subprocess
import torch
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
# check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
try:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sdxl"
except:
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sd15"
print(f"Loaded {sd_model_version} model!")
pipe = pipe.to(device, dtype=weight_dtype)
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
vae = pipe.vae
unet = pipe.unet
tokenizer_one = pipe.tokenizer
text_encoder_one = pipe.text_encoder
vae.requires_grad_(False)
if keep_vae_float32:
vae.to(device, dtype=torch.float32)
else:
vae.to(device, dtype=weight_dtype)
if weight_dtype != torch.float32:
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but may not be ideal for training..?")
unet.to(device, dtype=weight_dtype)
text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype)
tokenizer_two = text_encoder_two = None
if sd_model_version == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
return (
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version
def download_weights(url, dest):
start = time.time()
print("downloading url: ", url)
print("downloading to: ", dest, '...')
# Make sure the destination directory exists
dest_dir = os.path.dirname(dest)
if not os.path.exists(dest_dir):
os.makedirs(dest_dir)
try:
subprocess.check_call(["wget", "-q", "-O", dest, url])
except subprocess.CalledProcessError as e:
print("Error occurred while downloading:")
print("Exit status:", e.returncode)
print("Output:", e.output)
except Exception as e:
print("An unexpected error occurred:", e)
print(f"Downloading {url} took {time.time() - start} seconds")
def print_trainable_parameters(model, model_name = ''):
trainable_params = 0
all_param = 0
for name, param in model.named_parameters():
all_param += param.numel()
if param.requires_grad and "token_embedding" not in name:
trainable_params += param.numel()
line_delimiter = "#" * 80
print(line_delimiter)
print(
f"Trainable {model_name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
)
print(line_delimiter)
+276
View File
@@ -0,0 +1,276 @@
from peft import LoraConfig, get_peft_model
import torch
import prodigyopt
from typing import Iterable
def get_unet_optimizer(
prodigy_d_coef: float,
prodigy_growth_factor: float,
lora_weight_decay: float,
use_dora: bool,
unet_trainable_params: Iterable,
optimizer_name="prodigy"
):
## unet_trainable_params can be unet.parameters() or a list of lora params
# These learning rates will get overwritten in main.py:
if optimizer_name == "adamw":
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
elif optimizer_name == "AdamW8bit":
import bitsandbytes as bnb
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
elif optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_unet = prodigyopt.Prodigy(
unet_trainable_params,
d_coef = prodigy_d_coef,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=lora_weight_decay if not use_dora else 0.0,
betas=(0.9, 0.99),
growth_rate=prodigy_growth_factor # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
)
else:
raise NotImplementedError(f"Invalid optimizer_name for unet: {optimizer_name}")
print(f"Created {optimizer_name} optimizer for unet!")
return optimizer_unet
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
def is_belong_to_blocks(key, blocks):
try:
for g in blocks:
if g in key:
return True
return False
except Exception as e:
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
if use_blora:
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
try:
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
attns = [
attn_processor_name.rsplit(".", 1)[0]
for attn_processor_name, _ in unet.attn_processors.items()
if is_belong_to_blocks(attn_processor_name, blocks)
]
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
return target_modules
except Exception as e:
raise type(e)(
f"failed to get_target_modules, due to: {e}. "
f"Please check the modules specified in --lora_unet_blocks are correct"
)
def get_unet_lora_parameters(
lora_rank,
lora_alpha_multiplier: float,
lora_weight_decay: float,
use_dora: bool,
unet,
pipe,
):
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
unet_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian",
target_modules=target_modules,
use_dora=use_dora,
)
#unet.add_adapter(unet_lora_config)
unet = get_peft_model(unet, unet_lora_config)
pipe.unet = unet
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
unet_trainable_params = [
{
"params": unet_lora_parameters,
"weight_decay": lora_weight_decay if not use_dora else 0.0,
},
]
return unet, unet_trainable_params, unet_lora_parameters
def get_textual_inversion_optimizer(
text_encoders: list,
textual_inversion_lr: float,
textual_inversion_weight_decay,
optimizer_name: str
):
text_encoder_parameters = []
for text_encoder in text_encoders:
if text_encoder is not None:
text_encoder.train()
for name, param in text_encoder.named_parameters():
if "token_embedding" in name:
#param.data = param.to(dtype=torch.float32)
param.requires_grad = True
text_encoder_parameters.append(param)
print(f"Added {name} with shape {param.shape} to the trainable parameters")
else:
pass
params_to_optimize_ti = [
{
"params": text_encoder_parameters,
"lr": textual_inversion_lr if (optimizer_name != "prodigy") else 1.0,
"weight_decay":textual_inversion_weight_decay,
},
]
if optimizer_name == "prodigy":
optimizer_ti = prodigyopt.Prodigy(
params_to_optimize_ti,
d_coef = 1.0,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=textual_inversion_weight_decay,
betas=(0.9, 0.99),
#growth_rate=1.5, # this slows down the lr_rampup
)
elif optimizer_name == "adamw":
optimizer_ti = torch.optim.AdamW(
params_to_optimize_ti,
weight_decay=textual_inversion_weight_decay,
)
else:
raise NotImplementedError(f"Invalid optimizer_name: '{optimizer_name}'")
print(f"Created {optimizer_name} optimizer for textual inversion!")
return optimizer_ti, text_encoder_parameters
def get_text_encoder_lora_parameters(text_encoder, lora_rank, lora_alpha_multiplier, use_dora: bool):
text_encoder_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian",
target_modules=["k_proj", "q_proj", "v_proj", "out_proj"],
use_dora=use_dora,
)
text_encoder_peft_model = get_peft_model(text_encoder, text_encoder_lora_config)
text_encoder_lora_params = list(filter(lambda p: p.requires_grad, text_encoder_peft_model.parameters()))
return text_encoder_peft_model, text_encoder_lora_params
def get_optimizer_and_peft_models_text_encoder_lora(
text_encoders: list,
lora_rank: int,
lora_alpha_multiplier: float,
use_dora: bool,
optimizer_name: str,
lora_lr: float,
weight_decay: float
):
text_encoder_lora_parameters = []
text_encoder_peft_models = []
for text_encoder in text_encoders:
if text_encoder is not None:
text_encoder_peft_model, text_encoder_lora_params = get_text_encoder_lora_parameters(
text_encoder=text_encoder,
lora_rank=lora_rank,
lora_alpha_multiplier=lora_alpha_multiplier,
use_dora=use_dora
)
text_encoder_lora_parameters.extend(text_encoder_lora_params)
text_encoder_peft_models.append(text_encoder_peft_model)
else:
text_encoder_peft_models.append(None)
if optimizer_name == "adamw":
optimizer_text_encoder_lora = torch.optim.AdamW(
text_encoder_lora_parameters,
lr = lora_lr,
weight_decay=weight_decay if not use_dora else 0.0
)
else:
raise NotImplementedError(f"Text encoder LoRA finetuning is not yet implemented for optimizer: {optimizer_name}")
return optimizer_text_encoder_lora, text_encoder_peft_models
def get_current_lr(optimizer):
"""
Helper class to get the current lr for various types of optimizers
"""
try:
# Calculate the weighted average effective learning rate
total_lr = 0
total_params = 0
for group in optimizer.param_groups:
d = group['d']
lr = group['lr']
bias_correction = 1 # Default value
if group['use_bias_correction']:
beta1, beta2 = group['betas']
k = group['k']
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
effective_lr = d * lr * bias_correction
# Count the number of parameters in this group
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
total_lr += effective_lr * num_params
total_params += num_params
if total_params == 0:
return 0.0
else: return total_lr / total_params
except:
return optimizer.param_groups[0]['lr']
class OptimizerCollection:
def __init__(
self,
optimizer_textual_inversion = None,
optimizer_text_encoders = None,
optimizer_unet = None,
debug = False,
):
"""
run operations on all the relevant optimizers with a single function call
"""
self.debug = debug
self.optimizers = {
'textual_inversion': optimizer_textual_inversion,
'text_encoders': optimizer_text_encoders,
'unet': optimizer_unet
}
self.learning_rate_tracker = {'textual_inversion':[], 'text_encoders':[], 'unet':[]}
print("--> Initialized optimizers for:")
for key in self.optimizers.keys():
if self.optimizers[key] is not None:
print(key)
def get_lr(self, key):
return get_current_lr(self.optimizers[key])
def zero_grad(self):
for key in self.optimizers.keys():
if self.optimizers[key] is not None:
self.optimizers[key].zero_grad()
def step(self):
for key in self.optimizers.keys():
if self.optimizers[key] is not None:
self.optimizers[key].step()
if self.debug:
self.learning_rate_tracker[key].append(get_current_lr(self.optimizers[key]))
+399 -196
View File
@@ -1,7 +1,3 @@
# Have SwinIR upsample
# Have BLIP auto caption
# Have CLIPSeg auto mask concept
import gc
import fnmatch
import mimetypes
@@ -25,6 +21,7 @@ import numpy as np
import pandas as pd
import torch
from tqdm import tqdm
from transformers import (
BlipForConditionalGeneration,
Blip2ForConditionalGeneration,
@@ -36,14 +33,15 @@ from transformers import (
Swin2SRImageProcessor,
)
from io_utils import download_and_prep_training_data
from trainer.utils.io import download_and_prep_training_data
from trainer.utils.utils import fix_prompt
from trainer.config import model_paths
import re
import openai
from openai import OpenAI
from dotenv import load_dotenv
load_dotenv()
try:
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
client = OpenAI(api_key=OPENAI_API_KEY)
@@ -53,30 +51,20 @@ except:
client = None
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
MODEL_PATH = "./cache"
MAX_GPT_PROMPTS = 40
import re
def fix_prompt(prompt: str):
# Remove extra commas and spaces, and fix space before punctuation
prompt = re.sub(r"\s+", " ", prompt) # Replace multiple spaces with a single space
prompt = re.sub(r",,", ",", prompt) # Replace double commas with a single comma
prompt = re.sub(r"\s?,\s?", ", ", prompt) # Fix spaces around commas
prompt = re.sub(r"\s?\.\s?", ". ", prompt) # Fix spaces around periods
return prompt.strip() # Remove leading and trailing whitespace
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...)
MIN_GPT_PROMPTS = 3
MAX_GPT_PROMPTS = 50
def _find_files(pattern, dir="."):
"""Return list of files matching pattern in a given directory, in absolute format.
Unlike glob, this is case-insensitive.
"""
rule = re.compile(fnmatch.translate(pattern), re.IGNORECASE)
return [os.path.join(dir, f) for f in os.listdir(dir) if rule.match(f)]
def preprocess(
config,
working_directory,
concept_mode,
input_zip_path: Path,
@@ -85,7 +73,6 @@ def preprocess(
target_size: int,
crop_based_on_salience: bool,
use_face_detection_instead: bool,
temp: float,
left_right_flip_augmentation: bool = False,
augment_imgs_up_to_n: int = 0,
caption_model: str = "blip",
@@ -93,7 +80,6 @@ def preprocess(
) -> Path:
if os.path.exists(working_directory):
print(f"working_directory {working_directory} already existed.. deleting and recreating!")
shutil.rmtree(working_directory)
os.makedirs(working_directory)
@@ -108,7 +94,8 @@ def preprocess(
download_and_prep_training_data(input_zip_path, TEMP_IN_DIR)
n_training_imgs, trigger_text, segmentation_prompt, captions = load_and_save_masks_and_captions(
config = load_and_save_masks_and_captions(
config,
concept_mode,
files=TEMP_IN_DIR,
output_dir=TEMP_OUT_DIR,
@@ -118,13 +105,12 @@ def preprocess(
target_size=target_size,
crop_based_on_salience=crop_based_on_salience,
use_face_detection_instead=use_face_detection_instead,
temp=temp,
add_lr_flips = left_right_flip_augmentation,
augment_imgs_up_to_n = augment_imgs_up_to_n,
caption_model = caption_model
)
return Path(TEMP_OUT_DIR), n_training_imgs, trigger_text, segmentation_prompt, captions
return config, Path(TEMP_OUT_DIR)
@torch.no_grad()
@@ -148,7 +134,7 @@ def swin_ir_sr(
"""
model = Swin2SRForImageSuperResolution.from_pretrained(
model_id, cache_dir=MODEL_PATH
model_id, cache_dir = model_paths.get_path("SR")
).to(device)
processor = Swin2SRImageProcessor()
@@ -196,14 +182,15 @@ def clipseg_mask_generator(
if isinstance(target_prompts, str):
print(
f'Warning: only one target prompt "{target_prompts}" was given, so it will be used for all images'
f'Using "{target_prompts}" as CLIP-segmentation prompt for all images.'
)
target_prompts = [target_prompts] * len(images)
model = None
if any(target_prompts):
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
model = CLIPSegForImageSegmentation.from_pretrained(
model_id, cache_dir=MODEL_PATH
model_id, cache_dir = model_paths.get_path("CLIP")
).to(device)
masks = []
@@ -237,79 +224,68 @@ def clipseg_mask_generator(
masks.append(mask)
# cleanup
del model
gc.collect()
torch.cuda.empty_cache()
return masks
import textwrap
def cleanup_prompts_with_chatgpt(
prompts,
concept_mode, # face / object / style
seed, # seed for chatgpt reproducibility
verbose = True):
if concept_mode == "object_injection":
chat_gpt_prompt_1 = """
I have a set of images, each containing the same concept / figure. I have the following (poor) descriptions for each image:
"""
chat_gpt_prompt_2 = """
I want you to:
1. Find a good, short name/description of the single central concept that's in all the images. This [Concept Name] might eg already be present in the descriptions above, pick the most obvious name or words that would fit in all descriptions.
2. Insert the text "TOK, [Concept Name]" into all the descriptions above by rephrasing them where needed to naturally contain the text TOK, [Concept Name] while keeping as much of the description as possible.
Reply by first stating the "Concept Name:", followed by an enumerated list (using "-") of all the revised "Descriptions:".
"""
if concept_mode == "object":
chat_gpt_prompt_1 = """
Analyze a set of (poor) image descriptions each featuring the same concept, figure or thing.
Tasks:
1. Deduce a concise, fitting name for the concept that is visually descriptive (Concept Name).
2. Substitute the concept in each description with the placeholder "TOK", rearranging or adjusting the text where needed. Hallucinate TOK into the description if necessary (but dont mention when doing so, simply provide the final description)!
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of the placeholder string "TOK".
The descriptions are:
"""
chat_gpt_prompt_1 = textwrap.dedent("""
Analyze a set of (poor) image descriptions each featuring the same concept, figure or thing.
Tasks:
1. Deduce a concise (max 10 words) visual description of just the concept TOK (Concept Description), try to be as visually descriptive of TOK as possible!
2. Substitute the concept in each description with the placeholder "TOK", rearranging or adjusting the text where needed. Hallucinate TOK into the description if necessary (but dont mention when doing so, simply provide the final description)!
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of the placeholder string "TOK".
The descriptions are:""")
chat_gpt_prompt_2 = """
Respond with the chosen "Concept Name:" followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
"""
chat_gpt_prompt_2 = textwrap.dedent("""
Respond with "Concept Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
""")
elif concept_mode == "face":
chat_gpt_prompt_1 = """
Analyze a set of (poor) image descriptions, each featuring a person named TOK.
Tasks:
1. Rewrite each description, ensuring it refers only to a single person or character.
2. Integrate "a photo of TOK" naturally into each description, rearranging or adjusting where needed.
3. Streamline each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
The descriptions are:
"""
chat_gpt_prompt_1 = textwrap.dedent("""
Analyze a set of (poor) image descriptions, each featuring a person named TOK.
Tasks:
1. Deduce a concise (max 10 words) visual description of TOK (TOK Description), try to be as visually descriptive of TOK as possible, always mention their skin color, hallucinate a basic description if necessary (eg black man with long beard).
2. Rewrite each description, injecting "TOK" naturally into each description, adjusting where needed.
3. Streamline each description to focus on the context and surroundings of TOK instead of the visual appearance of TOK's face. Ensure mandatory inclusion of "TOK".
The descriptions are:""")
chat_gpt_prompt_2 = """
Respond with "Concept Name: TOK" followed by a list (using "-") of all the revised descriptions, each mentioning "a photo of TOK".
"""
chat_gpt_prompt_2 = textwrap.dedent("""
Respond with "TOK Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "TOK".
""")
elif concept_mode == "style":
chat_gpt_prompt_1 = """
Analyze a set of (poor) image descriptions, each featuring the same style named TOK.
Tasks:
1. Rewrite each description to focus solely on the TOK style.
2. Integrate "in the style of TOK" naturally into each description, typically at the beginning.
3. Summarize each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
The descriptions are:
"""
chat_gpt_prompt_1 = textwrap.dedent("""
Analyze a set of (poor) image descriptions, each featuring an example of a common aesthetic style named TOK.
Tasks:
1. Deduce a concise (max 7 words) visual description of the aesthetic style (Style Description).
2. Rewrite each description to focus solely on the non-stylistic contents of the image like characters, objects, colors, scene, context etc but not the stylistic elements captured by TOK.
3. Integrate "in the style of TOK" naturally into each description, typically at the beginning while summarizing each description to its core elements, ensuring clarity and mandatory inclusion of "TOK".
The descriptions are:""")
chat_gpt_prompt_2 = """
Respond with "Style Name: TOK" followed by a list (using "-") of all the revised descriptions, each mentioning "in the style of TOK".
"""
chat_gpt_prompt_2 = textwrap.dedent("""
Respond with "Style Description: ..." followed by a list (using "-") of all the revised descriptions, each mentioning "in the style of TOK".
""")
final_chatgpt_prompt = chat_gpt_prompt_1 + "\n- " + "\n- ".join(prompts) + "\n\n" + chat_gpt_prompt_2
final_chatgpt_prompt = chat_gpt_prompt_1 + "\n- " + "\n- ".join(prompts) + "\n" + chat_gpt_prompt_2
print("Final chatgpt prompt:")
print(final_chatgpt_prompt)
print("--------------------------")
print(f"Calling chatgpt with seed {seed}...")
response = client.chat.completions.create(
model="gpt-4-1106-preview",
model="gpt-4o",
seed=seed,
messages=[
{"role": "system", "content": "You are a helpful assistant."},
@@ -326,67 +302,70 @@ def cleanup_prompts_with_chatgpt(
# extract the final rephrased prompts from the response:
prompts = []
for line in gpt_completion.split("\n"):
if line.startswith("-"):
if line.startswith("-") or re.match(r'^\d+\.', line):
prompts.append(line[2:])
gpt_concept_name = extract_gpt_concept_name(gpt_completion, concept_mode)
trigger_text = "TOK, " + gpt_concept_name if concept_mode == 'object_injection' else "TOK"
gpt_concept_description = extract_gpt_concept_description(gpt_completion, concept_mode)
trigger_text = "TOK"
if concept_mode == 'style':
trigger_text = ", in the style of TOK"
gpt_concept_name = "" # Disables segmentation for style (use full img)
trigger_text = "in the style of TOK, "
return prompts, gpt_concept_name, trigger_text
return prompts, gpt_concept_description, trigger_text
def extract_gpt_concept_name(gpt_completion, concept_mode):
def extract_gpt_concept_description(gpt_completion, concept_mode):
"""
Extracts the concept name from the GPT completion based on the concept mode.
"""
concept_name = ""
prefix = ""
if concept_mode in ['face', 'style']:
concept_name = concept_mode
prefix = "Style Name:" if concept_mode == 'style' else ""
elif concept_mode in ['object_injection', 'object']:
prefix = "Concept Name:"
concept_mode = 'object_injection'
if prefix:
for line in gpt_completion.split("\n"):
if line.startswith(prefix):
concept_name = line[len(prefix):].strip()
break
if concept_mode == 'face':
prefix = "TOK Description:"
elif concept_mode == 'style':
prefix = "Style Description:"
elif concept_mode == 'object':
prefix = "Concept Description:"
for line in gpt_completion.split("\n"):
if line.startswith(prefix):
concept_name = line[len(prefix):].strip()
break
return concept_name
def post_process_captions(captions, text, concept_mode, job_seed):
text = text.strip()
print(f"Input captioning text: {text}")
gpt_cleanup_worked = False
gpt_concept_description = None
if len(captions) > 3 and len(captions) < MAX_GPT_PROMPTS and not text and client:
if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client:
retry_count = 0
while retry_count < 10:
while retry_count < 5:
try:
gpt_captions, gpt_concept_name, trigger_text = cleanup_prompts_with_chatgpt(captions, concept_mode, job_seed + retry_count)
gpt_captions, gpt_concept_description, trigger_text = cleanup_prompts_with_chatgpt(captions, concept_mode, job_seed + retry_count)
n_toks = sum("TOK" in caption for caption in gpt_captions)
if n_toks > int(0.8 * len(captions)) and (len(gpt_captions) == len(captions)):
# Ensure every caption contains "TOK"
# gpt-cleanup (mostly) worked, lets just ensure every caption contains "TOK" and finish
print("Making sure TOK is added to every training prompt...")
gpt_captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in gpt_captions]
captions = gpt_captions
gpt_cleanup_worked = True
break
else:
if len(gpt_captions) == len(captions):
print(f'GPT-4 did not return enough {n_toks}/{len(captions)} prompts containing "TOK", retrying...')
else:
print(f'GPT-4 returned the wrong number of prompts {len(gpt_captions)} instead of {len(captions)}, retrying...')
retry_count += 1
gpt_cleanup_worked = False
except Exception as e:
retry_count += 1
gpt_cleanup_worked = False
print(f"An error occurred after try {retry_count}: {e}")
time.sleep(1)
else:
gpt_concept_name, trigger_text = None, "TOK"
else:
time.sleep(0.5)
if not gpt_cleanup_worked:
# simple concat of trigger text with rest of prompt:
if len(text) == 0:
print("WARNING: no captioning text was given and we're not doing chatgpt cleanup...")
@@ -395,16 +374,14 @@ def post_process_captions(captions, text, concept_mode, job_seed):
trigger_text = "in the style of TOK, "
captions = [trigger_text + caption for caption in captions]
else:
trigger_text = "a photo of TOK, "
trigger_text = "TOK, "
captions = [trigger_text + caption for caption in captions]
else:
trigger_text = text
captions = [trigger_text + ", " + caption for caption in captions]
gpt_concept_name = None
captions = [fix_prompt(caption) for caption in captions]
return captions, trigger_text, gpt_concept_name
return captions, trigger_text, gpt_concept_description
def blip_caption_dataset(
@@ -416,19 +393,24 @@ def blip_caption_dataset(
"Salesforce/blip2-opt-2.7b",
] = "Salesforce/blip-image-captioning-large"
):
# If non of the captions are None, we dont need to do anything:
if all(captions):
print(f"All captions are already generated, skipping captioning...")
return captions
print(f"Using model {model_id} for image captioning...")
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "blip2" in model_id:
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = Blip2ForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device)
else:
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = BlipForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device)
for i, image in enumerate(tqdm(images)):
@@ -437,6 +419,10 @@ def blip_caption_dataset(
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65)
captions[i] = processor.decode(out[0], skip_special_tokens=True)
del model
gc.collect()
torch.cuda.empty_cache()
return captions
def encode_image(image_path):
@@ -454,6 +440,51 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
os.remove(output_path)
return base64_image
def gpt4_v_get_description(config, images):
if config.concept_mode == "object":
description = "object"
prompt = "Give a concise visual descriptioni of the object/figure/thing that all the grid-images have in common with at most 10 words. Dont start with statements like 'The image features...', just describe what you see."
elif config.concept_mode == "face":
description = "face"
prompt = "All the grid images depict a single person. Visually describe this person with at most 10 words. Dont start with statements like 'The image features...', just describe what you see. (eg an asian woman with long black hair)"
elif config.concept_mode == "style":
description = ""
prompt = "All these images share a common aesthetic style. Describe this style with at most 7 words. Dont start with statements like 'The image features...', just describe what you see. (eg impressionism collage surrealism)"
if not OPENAI_API_KEY:
print(f"Skipping GPT-4 Vision description because OPENAI_API_KEY is not set.")
return description
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {OPENAI_API_KEY}"
}
# TODO sample a grid img:
# .... TODO
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
payload = {
"model": "gpt-4o",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{base64_image}", "detail": "high"}}
]
}
],
"max_tokens": 60
}
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
answer = response.json()["choices"][0]["message"]["content"]
return captions
def gpt4_v_caption_dataset(
images, captions,
batch_size=4,
@@ -463,7 +494,7 @@ def gpt4_v_caption_dataset(
print(f"Skipping GPT-4 Vision captioning because OPENAI_API_KEY is not set.")
return captions
prompt = "Accurate describe the contents of this image without assumptions. Avoid starting with statements like 'The image features...', just describe what you see."
prompt = "Concisely describe this image without assumptions with at most 20 words. Dont start with statements like 'The image features...', just describe what you see."
headers = {
"Content-Type": "application/json",
@@ -474,7 +505,7 @@ def gpt4_v_caption_dataset(
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
payload = {
"model": "gpt-4-vision-preview",
"model": "gpt-4o",
"messages": [
{
"role": "user",
@@ -484,11 +515,18 @@ def gpt4_v_caption_dataset(
]
}
],
"max_tokens": 100
"max_tokens": 60
}
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
return index, response.json()["choices"][0]["message"]["content"]
try:
result = response.json()["choices"][0]["message"]["content"]
except:
print(response.json())
result = ""
return index, result
with concurrent.futures.ThreadPoolExecutor(max_workers=batch_size) as executor:
future_to_index = {executor.submit(fetch_caption, i, img): i for i, img in enumerate(images) if captions[i] is None}
@@ -519,49 +557,6 @@ def caption_dataset(
return captions
def _crop_to_square(
image: Image.Image, com: List[Tuple[int, int]], resize_to: Optional[int] = None
):
cx, cy = com
width, height = image.size
if width > height:
left_possible = max(cx - height / 2, 0)
left = min(left_possible, width - height)
right = left + height
top = 0
bottom = height
else:
left = 0
right = width
top_possible = max(cy - width / 2, 0)
top = min(top_possible, height - width)
bottom = top + width
image = image.crop((left, top, right, bottom))
if resize_to:
image = image.resize((resize_to, resize_to), Image.Resampling.LANCZOS)
return image
def _center_of_mass(mask: Image.Image):
"""
Returns the center of mass of the mask
"""
x, y = np.meshgrid(np.arange(mask.size[0]), np.arange(mask.size[1]))
mask_np = np.array(mask) + 0.01
x_ = x * mask_np
y_ = y * mask_np
x = np.sum(x_) / np.sum(mask_np)
y = np.sum(y_) / np.sum(mask_np)
return x, y
def load_image_with_orientation(path, mode = "RGB"):
image = Image.open(path)
@@ -640,9 +635,53 @@ def augment_image(image):
image = gaussian_blur(image)
return image
def round_to_nearest_multiple(x, multiple):
return int(float(multiple) * round(float(x) / float(multiple)))
'''
For Stable Diffusion 1.5, outputs are optimised around 512x512 pixels. Many common fine-tuned versions of SD1.5 are optimised around 768x768. The best resolutions for common aspect ratios are typically:
1:1 (square): 512x512, 768x768
3:2 (landscape): 768x512
2:3 (portrait): 512x768
4:3 (landscape): 768x576
3:4 (portrait): 576x768
16:9 (widescreen): 912x512
9:16 (tall): 512x912
For SDXL, outputs are optimised around 1024x1024 pixels. The best resolutions for common aspect ratios are typically:
stable-diffusion-xl-1024-v0-9 supports generating images at the following dimensions:
1024 x 1024
1152 x 896
896 x 1152
1216 x 832
832 x 1216
1344 x 768
768 x 1344
1536 x 640
640 x 1536
'''
def calculate_new_dimensions(target_size, target_aspect_ratio):
"""
Calculate the new width and height given a target size and aspect ratio.
"""
# Calculate the total number of pixels
n_pixels = target_size ** 2
# Calculate the new width and height based on the target aspect ratio
new_width = (n_pixels * target_aspect_ratio) ** 0.5
new_height = (n_pixels / new_width)
# round up/down to the nearest multiple of 64:
new_width = round_to_nearest_multiple(new_width, 64)
new_height = round_to_nearest_multiple(new_height, 64)
return [new_width, new_height]
def load_and_save_masks_and_captions(
config,
concept_mode: str,
files: Union[str, List[str]],
output_dir: str = "tmp_out",
@@ -652,7 +691,6 @@ def load_and_save_masks_and_captions(
target_size: int = 1024,
crop_based_on_salience: bool = True,
use_face_detection_instead: bool = False,
temp: float = 1.0,
n_length: int = -1,
add_lr_flips: bool = False,
augment_imgs_up_to_n: int = 0,
@@ -668,7 +706,6 @@ def load_and_save_masks_and_captions(
# load images
if isinstance(files, str):
if os.path.isdir(files):
print("Scanning directory for images...")
files = (
_find_files("*.png", files)
+ _find_files("*.jpg", files)
@@ -677,7 +714,7 @@ def load_and_save_masks_and_captions(
if len(files) == 0:
raise Exception(
f"No files found in {files}. Either {files} is not a directory or it does not contain any .png or .jpg/jpeg files."
f"No images were found... Are you sure you provided a valid dataset?"
)
if n_length == -1:
n_length = len(files)
@@ -693,6 +730,30 @@ def load_and_save_masks_and_captions(
else:
captions.append(None)
# Compute average aspect ratio of images:
aspect_ratios = [image.size[0] / image.size[1] for image in images]
avg_aspect_ratio = sum(aspect_ratios) / len(aspect_ratios)
print(f"Average aspect ratio of images (width / height): {avg_aspect_ratio:.3f}")
config.train_img_size = calculate_new_dimensions(target_size, avg_aspect_ratio)
config.train_aspect_ratio = config.train_img_size[0] / config.train_img_size[1]
target_size = max(config.train_img_size)
print(f"New train_img_size: {config.train_img_size}")
if config.validation_img_size is None:
config.validation_img_size = [0, 0]
multiplier = 2.0 if config.sd_model_version == "sdxl" else 1.0
config.validation_img_size[0] = config.train_img_size[0] * multiplier
config.validation_img_size[1] = config.train_img_size[1] * multiplier
elif isinstance(config.validation_img_size, int):
n_pixels = config.validation_img_size ** 2
config.validation_img_size = [0, 0]
config.validation_img_size[0] = (n_pixels * config.train_aspect_ratio) ** 0.5
config.validation_img_size[1] = (n_pixels / config.validation_img_size[0])
config.validation_img_size[0] = round_to_nearest_multiple(config.validation_img_size[0], 64)
config.validation_img_size[1] = round_to_nearest_multiple(config.validation_img_size[1], 64)
print(f"Validation_img_size was set to: {config.validation_img_size}")
n_training_imgs = len(images)
n_captions = len([c for c in captions if c is not None])
print(f"Loaded {n_training_imgs} images, {n_captions} of which have captions.")
@@ -700,27 +761,47 @@ def load_and_save_masks_and_captions(
if len(images) < 50: # upscale images that are smaller than target_size:
print("upscaling imgs..")
upscale_margin = 0.75
images = swin_ir_sr(images, target_size=(int(target_size*upscale_margin), int(target_size*upscale_margin)))
images = swin_ir_sr(images, target_size=(int(config.train_img_size[0]*upscale_margin), int(config.train_img_size[0]*upscale_margin)))
if add_lr_flips and len(images) < 40:
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
captions = captions + captions
# It's nice if we can achieve the gpt pass, so pre-augment the images if there's very few:
# Ensure we have at least 'augment_imgs_up_to_n' images through augmentation
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < MIN_GPT_PROMPTS:
print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images])
aug_caps.extend(captions)
images.extend(aug_imgs)
captions.extend(aug_caps)
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
images = images[:MAX_GPT_PROMPTS-1]
captions = captions[:MAX_GPT_PROMPTS-1]
# Use BLIP for autocaptioning:
if len(images) > 50 and caption_model != "blip":
print(f"Captioning a lot of ({len(images)}) images --> falling back to using blip!")
caption_model = "blip"
print(f"Generating {len(images)} captions using mode: {concept_mode}...")
captions = caption_dataset(images, captions, caption_model = caption_model)
# Cleanup prompts using chatgpt:
captions, trigger_text, gpt_concept_name = post_process_captions(captions, caption_text, concept_mode, seed)
captions = [fix_prompt(caption) for caption in captions]
trigger_text = ""
gpt_concept_description = None
if not config.disable_ti:
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
aug_imgs, aug_caps = [],[]
while len(images) + len(aug_imgs) < augment_imgs_up_to_n: # if we still have a very small amount of imgs, do some basic augmentation:
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images])
aug_caps.extend(captions)
@@ -728,25 +809,30 @@ def load_and_save_masks_and_captions(
images.extend(aug_imgs)
captions.extend(aug_caps)
if (gpt_concept_name is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_name}")
mask_target_prompts = gpt_concept_name
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
mask_target_prompts = gpt_concept_description
if mask_target_prompts is None:
if mask_target_prompts is None or config.concept_mode == "style":
print("Disabling CLIP-segmentation")
mask_target_prompts = ""
temp = 999
else:
temp = config.clipseg_temperature
print(f"Generating {len(images)} masks...")
# Make sure we have a bias for the background pixels to never 100% ignore them
background_bias = 0.05
if not use_face_detection_instead:
seg_masks = clipseg_mask_generator(
images=images, target_prompts=mask_target_prompts, temp=temp
images=images, target_prompts=mask_target_prompts, temp=temp, bias=background_bias
)
else:
mask_target_prompts = "FACE detection was used"
if add_lr_flips:
print("WARNING you are applying face detection while also doing left-right flips, this might not be what you intended?")
seg_masks = face_mask_google_mediapipe(images=images)
seg_masks = face_mask_google_mediapipe(images=images, bias=background_bias*255)
print("Masks generated! Cropping images to center of mass...")
# find the center of mass of the mask
@@ -756,23 +842,31 @@ def load_and_save_masks_and_captions(
coms = [(image.size[0] / 2, image.size[1] / 2) for image in images]
# based on the center of mass, crop the image to a square
print("Cropping squares...")
print("Cropping and resizing images...")
images = [
_crop_to_square(image, com, resize_to=None)
_crop_to_aspect_ratio(image, com, target_aspect_ratio = config.train_aspect_ratio, # width / height
resize_to = target_size)
for image, com in zip(images, coms)
]
seg_masks = [
_crop_to_square(mask, com, resize_to=target_size)
_crop_to_aspect_ratio(mask, com, target_aspect_ratio = config.train_aspect_ratio, # width / height
resize_to = target_size)
for mask, com in zip(seg_masks, coms)
]
print("Resizing images to training size...")
images = [
image.resize((target_size, target_size), Image.Resampling.LANCZOS)
for image in images
]
print("Expanding masks...")
if use_face_detection_instead:
dilation_radius = -0.02 * (config.train_img_size[0] + config.train_img_size[0]) / 2
blur_radius = 0.02 * (config.train_img_size[0] + config.train_img_size[0]) / 2
else:
dilation_radius = 0.0
blur_radius = 0.005 * (config.train_img_size[0] + config.train_img_size[0]) / 2
for i in range(len(seg_masks)):
seg_masks[i] = grow_mask(seg_masks[i], dilation_radius=dilation_radius, blur_radius=blur_radius)
print("Done!")
data = []
# clean TEMP_OUT_DIR first
if os.path.exists(output_dir):
@@ -780,9 +874,22 @@ def load_and_save_masks_and_captions(
os.remove(os.path.join(output_dir, file))
os.makedirs(output_dir, exist_ok=True)
if config.disable_ti:
print('------------------ WARNING -------------------')
print("Removing 'TOK, ' from captions...")
print("This will completely disable textual_inversion!!")
print('------------------ WARNING -------------------')
if gpt_concept_description:
replace_str = gpt_concept_description
else:
replace_str = ""
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
captions = [caption.replace("TOK", replace_str) for caption in captions]
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
# Make sure we've correctly inserted the TOK into every caption:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
print("Final captions:")
for caption in captions:
print(caption)
@@ -806,12 +913,111 @@ def load_and_save_masks_and_captions(
df.to_csv(os.path.join(output_dir, "captions.csv"), index=False)
print("---> Training data 100% ready to go!")
return n_training_imgs, trigger_text, mask_target_prompts, captions
# do a final prompt cleaning pass to fix weird commas and spaces:
captions = [fix_prompt(caption) for caption in captions]
# Update the training attributes with some info from the pre-processing:
config.training_attributes["n_training_imgs"] = n_training_imgs
config.training_attributes["trigger_text"] = trigger_text
config.training_attributes["segmentation_prompt"] = mask_target_prompts
config.training_attributes["gpt_description"] = gpt_concept_description
config.training_attributes["captions"] = captions
return config
from PIL import Image, ImageFilter, ImageChops
def grow_mask(mask, dilation_radius=5, blur_radius=3):
dilation_radius = int(dilation_radius)
blur_radius = int(blur_radius)
# Load the image
mask = mask.convert('L') # Ensure it's in grayscale
# Get the minimum pixel value in the mask:
min_mask_value = int(np.min(np.array(mask)))
# Dilate the mask
if dilation_radius > 0:
mask = mask.filter(ImageFilter.MinFilter(dilation_radius * 2 + 1))
# Apply Gaussian blur to the dilated mask
if blur_radius > 0:
mask = mask.filter(ImageFilter.GaussianBlur(blur_radius))
# Clip the mask pixel values to make sure they dont go below the minimum value
mask = ImageChops.lighter(mask, Image.new('L', mask.size, min_mask_value))
return mask
def _center_of_mass(mask: Image.Image):
"""
Returns the center of mass of the mask
"""
x, y = np.meshgrid(np.arange(mask.size[0]), np.arange(mask.size[1]))
mask_np = np.array(mask) + 0.01
x_ = x * mask_np
y_ = y * mask_np
x = np.sum(x_) / np.sum(mask_np)
y = np.sum(y_) / np.sum(mask_np)
return x, y
def _crop_to_aspect_ratio(
image: Image.Image,
com: List[Tuple[int, int]],
target_aspect_ratio: float = 1.0, # width / height
resize_to: Optional[int] = None
):
"""
Crops the image to the specified aspect ratio around the center of mass of the mask.
"""
cx, cy = com
width, height = image.size
if target_aspect_ratio > 1: # Wider than tall
new_width = int(min(width, height * target_aspect_ratio))
new_height = int(new_width / target_aspect_ratio)
else: # Taller than wide or square
new_height = int(min(height, width / target_aspect_ratio))
new_width = int(new_height * target_aspect_ratio)
left = int(max(cx - new_width / 2, 0))
right = int(min(left + new_width, width))
top = int(max(cy - new_height / 2, 0))
bottom = int(min(top + new_height, height))
# Adjust if the crop goes beyond the image boundaries
if right > width:
overshoot = right - width
right = width
left = max(0, left - overshoot) # Adjust left as well symmetrically
if bottom > height:
overshoot = bottom - height
bottom = height
top = max(0, top - overshoot) # Adjust top as well symmetrically
image = image.crop((left, top, right, bottom))
if resize_to:
if target_aspect_ratio > 1:
resize_height = int(resize_to / target_aspect_ratio)
image = image.resize((resize_to, resize_height), Image.Resampling.LANCZOS)
else:
resize_width = int(resize_to * target_aspect_ratio)
image = image.resize((resize_width, resize_to), Image.Resampling.LANCZOS)
return image
def face_mask_google_mediapipe(
images: List[Image.Image], blur_amount: float = 0.0, bias: float = 50.0
images: List[Image.Image], blur_amount: float = 0.0, bias: float = 10.0
) -> List[Image.Image]:
"""
Returns a list of images with masks on the face parts.
@@ -852,8 +1058,6 @@ def face_mask_google_mediapipe(
min(ih - bbox[1], bbox[3]),
)
print(bbox)
# Extract face landmarks
face_landmarks = face_mesh.process(
image_np[bbox[1] : bbox[1] + bbox[3], bbox[0] : bbox[0] + bbox[2]]
@@ -929,15 +1133,14 @@ def face_mask_google_mediapipe(
# Convert mask to 'L' mode (grayscale) before saving
mask = mask.convert("L")
masks.append(mask)
else:
# If face landmarks are not available, add a black mask of the same size as the image
masks.append(Image.new("L", (iw, ih), 255))
masks.append(Image.new("L", (iw, ih), 0))
else:
print("No face detected, adding full mask")
# If no face is detected, add a white mask of the same size as the image
masks.append(Image.new("L", (iw, ih), 255))
# If no face is detected, add a black mask of the same size as the image
masks.append(Image.new("L", (iw, ih), 0))
return masks
+289
View File
@@ -0,0 +1,289 @@
from functools import reduce
from diffusers import StableDiffusionXLPipeline
from diffusers.models.attention_processor import AttnProcessor2_0, Attention
from typing import Optional
import torch
import torch.nn as nn
from diffusers.utils.deprecation_utils import deprecate
import torch.nn.functional as F
import math
from einops.layers.torch import Reduce
from torchtyping import TensorType
from einops import rearrange
# Find all instances of AttnProcessor2_0 in the UNet
def find_attnprocessor2_0(unet):
module_names = []
"""
this function assumes that there are fewer than 50 down blocks, attention modules and transformer blocks
if you're not sure, feel free to set it to an arbitrarily large number
don't worry, it won't slow anything down.
"""
for block_type in ["down_blocks", "up_blocks"]:
for down_block_index in range(50):
for attentions_index in range(50):
for transformer_blocks_index in range(50):
example_module_name = f"{block_type}.{down_block_index}.attentions.{attentions_index}.transformer_blocks.{transformer_blocks_index}.attn2.processor"
try:
module = get_module_by_name(module=unet, name = example_module_name)
assert isinstance(module, AttnProcessor2_0), f"Expected module to be an instance of AttnProcessor2_0 but found it to be: {type(module)}"
# print(f"Found: {example_module_name}")
module_names.append(example_module_name)
except AttributeError:
# print(f"Ignored name: {example_module_name}\nsince it does not exist")
pass
print(f"Found: {len(module_names)} modules")
return module_names
class DAAMLossAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(self, name: str):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
self.name = name
self.cross_attention_scores = None
self.reduce_op = Reduce(
"batch heads img text -> batch img text",
reduction="sum"
)
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
temb: Optional[torch.Tensor] = None,
*args,
**kwargs,
) -> torch.Tensor:
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
deprecate("scale", "1.0.0", deprecation_message)
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
"""
Mayukh's experiment
"""
mayukh_experiment = False
if encoder_hidden_states is not None:
"""
this triggers cross attn
"""
mayukh_experiment = True
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
if mayukh_experiment:
# Calculate QK^T
qk_t = torch.matmul(query, key.transpose(-2, -1))
# Calculate attention scores (scaled QK^T)
d_k = query.size(-1) # Assuming the last dimension is the embedding dimension
attention_scores = qk_t / math.sqrt(d_k)
attention_scores = self.reduce_op(
attention_scores,
)
self.cross_attention_scores = attention_scores
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class DAAMLoss:
def __init__(self, attention_processors: list[DAAMLossAttnProcessor2_0]):
self.attention_processors = attention_processors
self.layer_names = [
x.name for x in attention_processors
]
def get_all_cross_attention_scores(self):
cross_attention_scores = {}
for p in self.attention_processors:
cross_attention_scores[
p.name
] = p.cross_attention_scores
return cross_attention_scores
def compute_single_token_loss(self, text_token_index: list[int], reduce = False):
loss = {}
cross_attention_scores = self.get_all_cross_attention_scores()
for name, cross_attention_map in cross_attention_scores.items():
"""
cross_attention_map.shape: (batch, image_patches, text_tokens)
"""
assert cross_attention_map.ndim == 3
loss[name] = cross_attention_map[:,:,text_token_index].norm() / cross_attention_map.shape[1]
if reduce:
all_losses = list(loss.values())
return sum(all_losses)/len(all_losses)
else:
return loss
def compute_loss(self, text_token_indices: list[int], reduce = False):
losses = []
for text_token_index in text_token_indices:
losses.append(
self.compute_single_token_loss(
text_token_index=text_token_index,
reduce = True
)
)
if reduce:
return sum(losses)/len(losses)
else:
return losses
def get_image_heatmap(self, text_token_index: int, layer_name: str) -> TensorType["batch", "height", "width"]:
cross_attention_scores = self.get_all_cross_attention_scores()
assert layer_name in list(cross_attention_scores.keys())
cross_attention_scores_single_token = cross_attention_scores[layer_name][:,:,text_token_index]
assert cross_attention_scores_single_token.ndim == 2 ## batch, hw
heatmap = rearrange(
cross_attention_scores_single_token,
"batch (height width) -> batch height width",
height = int(math.sqrt(cross_attention_scores_single_token.shape[1])),
width = int(math.sqrt(cross_attention_scores_single_token.shape[1]))
)
return heatmap
def get_the_daam_heatmap(self, text_token_index: int) ->TensorType["batch", "height", "width"]:
all_heatmaps = []
for layer_name in self.layer_names:
heatmap = self.get_image_heatmap(
text_token_index=text_token_index,
layer_name=layer_name
)
all_heatmaps.append(heatmap)
## each heatmap has a shape: batch, h, w where h=w
## now find the maximum possible height and width across all heatmaps
max_height = max(heatmap.shape[1] for heatmap in all_heatmaps)
max_width = max(heatmap.shape[2] for heatmap in all_heatmaps)
## now resize all_heatmaps to (batch, max_height, max_width) using F.interpolate
resized_heatmaps = [
F.interpolate(input = x.unsqueeze(1), size = (max_height, max_width)).squeeze(1)
for x in all_heatmaps
]
return sum(resized_heatmaps)
def get_module_by_name(module: nn.Module, name: str):
"""Retrieve a module nested in another by its access string."""
if name == "":
return module
names = name.split(sep=".")
return reduce(getattr, names, module)
def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionXLPipeline, DAAMLoss]:
assert isinstance(pipeline, StableDiffusionXLPipeline)
## find out where the attention processor thingies are
module_names = find_attnprocessor2_0(
unet = pipeline.unet
)
all_daam_attention_processors = []
# override the attention processor thingies
for name in module_names:
# print(f"Replacing: {name}")
# Get parent module and attribute name
parent_name = ".".join(name.split(".")[:-1])
attr_name = name.split(".")[-1]
# Get the parent module
parent_module = get_module_by_name(module=pipeline.unet, name=parent_name)
daam_attention_processor = DAAMLossAttnProcessor2_0(name=name)
all_daam_attention_processors.append(daam_attention_processor)
# Set the attribute
setattr(parent_module, attr_name, daam_attention_processor)
# Verify the replacement
current_module = get_module_by_name(module=pipeline.unet, name=name)
assert isinstance(current_module, DAAMLossAttnProcessor2_0)
daam_loss = DAAMLoss(
attention_processors=all_daam_attention_processors
)
return pipeline, daam_loss
-590
View File
@@ -1,590 +0,0 @@
import os
import math
import random
import numpy as np
import torch
import fnmatch
from peft import LoraConfig, get_peft_model
from diffusers.optimization import get_scheduler
from tqdm import tqdm
import shutil
import time
import gc
import prodigyopt
from .config import (
TrainerConfig,
precision_map
)
from .dataset_and_utils import (
load_models,
TokenEmbeddingsHandler,
PreprocessedDataset,
plot_torch_hist,
plot_loss,
plot_lrs
)
from .utils.model_info import print_trainable_parameters
from .utils.snr import compute_snr
from .utils.learning_rate import get_avg_lr
from .utils.lora import save_lora
from .utils.rendering import render_images
from io_utils import download_weights
from preprocess import preprocess
class Trainer:
def __init__(self, args):
self.args = args
random.seed(args.seed)
torch.manual_seed(args.seed)
np.random.seed(args.seed)
torch.cuda.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
#torch.backends.cudnn.deterministic = True
print("Trainer initialized!")
def train(self):
if self.args.concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
self.args.l1_penalty = 0.05
args = self.args
if args.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
weight_dtype = precision_map[args.precision]
print(f"Loading models with weight_dtype: {weight_dtype}")
if args.scale_lr_based_on_grad_acc:
unet_learning_rate = (
args.unet_learning_rate * args.gradient_accumulation_steps * args.train_batch_size
)
# Download the weights if they don't exist locally
if not os.path.exists(args.pretrained_model['path']):
download_weights(args.pretrained_model['url'], args.pretrained_model['path'])
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(
pretrained_model = args.pretrained_model,
device=args.device,
weight_dtype=weight_dtype
)
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
[text_encoder_one, text_encoder_two], [tokenizer_one, tokenizer_two]
)
starting_toks = None
embedding_handler.initialize_new_tokens(
inserting_toks=args.inserting_list_tokens,
starting_toks=starting_toks,
seed=args.seed
)
text_encoders = [text_encoder_one, text_encoder_two]
unet_param_to_optimize = []
text_encoder_parameters = []
for text_encoder in text_encoders:
if text_encoder is not None:
for name, param in text_encoder.named_parameters():
if "token_embedding" in name:
param.requires_grad = True
text_encoder_parameters.append(param)
else:
param.requires_grad = False
unet_param_to_optimize_names = []
unet_lora_parameters = []
if not args.is_lora:
WHITELIST_PATTERNS = [
# "*.attn*.weight",
# "*ff*.weight",
"*"
]
BLACKLIST_PATTERNS = ["*.norm*.weight", "*time*"]
for name, param in unet.named_parameters():
if any(
fnmatch.fnmatch(name, pattern) for pattern in WHITELIST_PATTERNS
) and not any(
fnmatch.fnmatch(name, pattern) for pattern in BLACKLIST_PATTERNS
):
param.requires_grad_(True)
unet_param_to_optimize_names.append(name)
print(f"Training: {name}")
else:
param.requires_grad_(False)
# Optimizer creation
params_to_optimize = [
{
"params": text_encoder_parameters,
"lr": args.textual_inversion_lr,
"weight_decay": args.textual_inversion_weight_decay,
},
]
params_to_optimize_prodigy = [
{
"params": unet_param_to_optimize,
"lr": unet_learning_rate,
"weight_decay": args.lora_weight_decay,
},
]
else:
# Do lora-training instead.
unet.requires_grad_(False)
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
use_dora = True
unet_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights="gaussian",
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
use_dora=use_dora,
)
if use_dora:
print(f"Disabling L1 penalty for DORA training")
args.l1_penalty = 0.0
unet = get_peft_model(unet, unet_lora_config)
print_trainable_parameters(unet, name = 'unet')
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
# Loop over the unet_lora_parameters and print their names and shapes:
for name, param in unet.named_parameters():
if param.requires_grad:
print(name, param.shape)
params_to_optimize = [{
"params": text_encoder_parameters,
"lr": args.textual_inversion_lr,
"weight_decay": args.textual_inversion_weight_decay,
}]
params_to_optimize_prodigy = [{
"params": unet_lora_parameters,
"lr": 1.0,
"weight_decay": args.lora_weight_decay,
}]
if args.optimizer_name == "adamw":
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=0.0, # this wd doesn't matter, I think
)
optimizer_prod = None
elif args.optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_prod = prodigyopt.Prodigy(
params_to_optimize_prodigy,
d_coef = args.prodigy_d_coef,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=args.lora_weight_decay,
betas=(0.9, 0.99),
growth_rate=1.025, # this slows down the lr_rampup
#growth_rate=1.05, # this slows down the lr_rampup
)
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=args.textual_inversion_weight_decay,
)
train_dataset = PreprocessedDataset(
args.instance_data_dir,
tokenizer_one,
tokenizer_two,
vae,
do_cache=args.train_dataset_cache,
substitute_caption_map=args.token_dict,
)
print(f"# PTI : Loaded dataset, do_cache: {args.train_dataset_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=args.train_batch_size,
shuffle=True,
num_workers=args.dataloader_num_workers,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps
)
if args.max_train_steps is None:
max_train_steps = num_train_epochs * num_update_steps_per_epoch
else:
max_train_steps = args.max_train_steps
lr_scheduler = get_scheduler(
args.lr_scheduler_name,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * args.gradient_accumulation_steps,
num_training_steps=max_train_steps * args.gradient_accumulation_steps,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps
)
num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch)
total_batch_size = args.train_batch_size * args.gradient_accumulation_steps
if args.verbose:
print(f"# PTI : Running training ")
print(f"# PTI : Num examples = {len(train_dataset)}")
print(f"# PTI : Num batches each epoch = {len(train_dataloader)}")
print(f"# PTI : Num Epochs = {num_train_epochs}")
print(f"# PTI : Instantaneous batch size per device = {args.train_batch_size}")
print(f"# PTI : Total train batch size (distributed & accumulation) = {total_batch_size}")
print(f"# PTI : Gradient Accumulation steps = {args.gradient_accumulation_steps}")
print(f"# PTI : Total optimization steps = {max_train_steps}")
global_step = 0
first_epoch = 0
last_save_step = 0
progress_bar = tqdm(range(global_step, max_train_steps), position=0, leave=True)
checkpoint_dir = os.path.join(args.output_dir, "checkpoints")
if os.path.exists(checkpoint_dir):
shutil.rmtree(checkpoint_dir)
os.makedirs(f"{checkpoint_dir}")
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
#embedding_handler.pre_optimize_token_embeddings(train_dataset)
ti_lrs, lora_lrs = [], []
losses = []
start_time, images_done = time.time(), 0
for epoch in range(first_epoch, num_train_epochs):
unet.train()
progress_bar.set_description(f"# PTI :step: {global_step}, epoch: {epoch}")
for step, batch in enumerate(train_dataloader):
progress_bar.update(1)
if args.hard_pivot:
if epoch >= num_train_epochs // 2:
if optimizer is not None:
print("----------------------")
print("# PTI : Pivot halfway")
print("----------------------")
# remove text encoder parameters from the optimizer
optimizer.param_groups = None
# remove the optimizer state corresponding to text_encoder_parameters
for param in text_encoder_parameters:
if param in optimizer.state:
del optimizer.state[param]
optimizer = None
else: # Update learning rates gradually:
finegrained_epoch = epoch + step / len(train_dataloader)
completion_f = finegrained_epoch / num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training
optimizer.param_groups[0]['lr'] = args.textual_inversion_lr * (1 - completion_f) ** 2.0
try: #sdxl
(tok1, tok2), vae_latent, mask = batch
except: #sd15
tok1, vae_latent, mask = batch
tok2 = None
vae_latent = vae_latent.to(weight_dtype)
# tokens to text embeds
prompt_embeds_list = []
for tok, text_encoder in zip((tok1, tok2), text_encoders):
if tok is None:
continue
prompt_embeds_out = text_encoder(
tok.to(text_encoder.device),
output_hidden_states=True,
)
pooled_prompt_embeds = prompt_embeds_out[0]
prompt_embeds = prompt_embeds_out.hidden_states[-2]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
prompt_embeds_list.append(prompt_embeds)
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
# Create Spatial-dimensional conditions.
original_size = (args.resolution, args.resolution)
target_size = (args.resolution, args.resolution)
crops_coords_top_left = (
args.crops_coords_top_left_h,
args.crops_coords_top_left_w
)
add_time_ids = list(original_size + crops_coords_top_left + target_size)
add_time_ids = torch.tensor([add_time_ids])
add_time_ids = add_time_ids.to(
args.device,
dtype=prompt_embeds.dtype
).repeat(
bs_embed, 1
)
# Sample noise that we'll add to the latents:
noise = torch.randn_like(vae_latent)
noise_offset = 0.05 # TODO, turn this into an input arg and do a grid search
if noise_offset > 0.0:
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
noise += noise_offset * torch.randn(
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
bsz = vae_latent.shape[0]
timesteps = torch.randint(
0,
noise_scheduler.config.num_train_timesteps,
(bsz,),
device=vae_latent.device,
).long()
noisy_model_input = noise_scheduler.add_noise(vae_latent, noise, timesteps)
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,1:-2,:]) * noise_sigma
# Predict the noise residual
model_pred = unet(
noisy_model_input,
timesteps,
prompt_embeds,
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
).sample
# Get the unet prediction target depending on the prediction type:
if noise_scheduler.config.prediction_type == "epsilon":
target = noise
else:
raise NotImplementedError(f"Not implemented for noise_scheduler.config.prediction_type: {noise_scheduler.config.prediction_type}")
# Compute the loss:
if args.snr_gamma is None:
loss = (model_pred - target).pow(2) * mask
# modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
# Average the normalized errors across the batch
loss = loss.mean()
else:
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
# This is discussed in Section 4.2 of the same paper.
snr = compute_snr(noise_scheduler, timesteps)
base_weight = (
torch.stack([snr, args.snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
)
if noise_scheduler.config.prediction_type == "v_prediction":
# Velocity objective needs to be floored to an SNR weight of one.
mse_loss_weights = base_weight + 1
else:
# Epsilon and sample both use the same loss weights.
mse_loss_weights = base_weight
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
loss = (model_pred - target).pow(2) * mask
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
if 1: # modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
loss = loss.mean()
if args.l1_penalty > 0.0:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += args.l1_penalty * l1_norm
# Print the relative L1 norm:
if global_step % 50 == 0:
print(f" ---- L1 norm: {l1_norm.item():.4f}")
print(f" ---- L1 loss: {args.l1_penalty * l1_norm.item():.4f}")
print(f" ---- Total loss: {loss.item():.4f}")
losses.append(loss.item())
loss = loss / args.gradient_accumulation_steps
loss.backward()
'''
apart from the usual gradient accumulation steps,
we also do a backward pass after computing the last forward pass in the epoch (last_batch == True)
this is to make sure that we're not missing out on any data
'''
last_batch = (step + 1 == len(train_dataloader))
if (step + 1) % args.gradient_accumulation_steps == 0 or last_batch:
if optimizer is not None:
optimizer.step()
optimizer.zero_grad()
if optimizer_prod is not None:
optimizer_prod.step()
optimizer_prod.zero_grad()
# after every optimizer step, we reset the non-trainable embeddings to the original embeddings
embedding_handler.retract_embeddings(print_stds = (global_step % 50 == 0))
embedding_handler.fix_embedding_std(args.off_ratio_power)
# Track the learning rates for final plotting:
lora_lrs.append(get_avg_lr(optimizer_prod))
try:
ti_lrs.append(optimizer.param_groups[0]['lr'])
except:
ti_lrs.append(0.0)
# Print some statistics:
if (global_step % args.checkpointing_steps == 0): # and (global_step > 0):
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
save_lora(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=args.token_dict,
args_dict=args.args_dict,
is_lora= args.is_lora,
unet_lora_parameters=unet_lora_parameters,
unet_param_to_optimize_names=unet_param_to_optimize_names
)
args.save_as_json(os.path.join(output_save_dir,"training_args.json"))
last_save_step = global_step
validation_prompts = render_images(
pipe, target_size,
output_save_dir,
global_step,
args.seed,
args.is_lora,
args.pretrained_model,
n_imgs = 4
)
if args.debug:
token_embeddings = embedding_handler.get_trainable_embeddings()
for i, token_embeddings_i in enumerate(token_embeddings):
plot_torch_hist(
token_embeddings_i[0],
global_step,
args.output_dir,
f"embeddings_weights_token_0_{i}",
min_val=-0.05,
max_val=0.05,
ymax_f = 0.05
)
plot_torch_hist(
token_embeddings_i[1],
global_step,
args.output_dir,
f"embeddings_weights_token_1_{i}",
min_val=-0.05,
max_val=0.05,
ymax_f = 0.05
)
embedding_handler.print_token_info()
plot_torch_hist(
unet_lora_parameters,
global_step,
args.output_dir,
"lora_weights",
min_val=-0.3,
max_val=0.3,
ymax_f = 0.05
)
plot_loss(losses, save_path=f'{args.output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{args.output_dir}/learning_rates.png')
gc.collect()
torch.cuda.empty_cache()
images_done += args.train_batch_size
global_step += 1
if global_step % 100 == 0:
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
if args.debug:
plot_loss(losses, save_path=f'{args.output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{args.output_dir}/learning_rates.png')
plot_torch_hist(unet_lora_parameters, global_step, args.output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
plot_torch_hist(embedding_handler.get_trainable_embeddings(), global_step, args.output_dir, "embeddings_weights", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
# final_save
if (global_step - last_save_step) > 51:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
else:
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
if not os.path.exists(output_save_dir):
save_lora(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=args.token_dict,
args_dict=args.args_dict,
is_lora= args.is_lora,
unet_lora_parameters=unet_lora_parameters,
unet_param_to_optimize_names=unet_param_to_optimize_names
)
args.save_as_json(os.path.join(output_save_dir,"training_args.json"))
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, args.seed, args.is_lora, args.pretrained_model, n_imgs = 4, n_steps = 35)
else:
print(f"Skipping final save, {output_save_dir} already exists")
del unet
del vae
del text_encoder_one
del text_encoder_two
del tokenizer_one
del tokenizer_two
del embedding_handler
del pipe
gc.collect()
torch.cuda.empty_cache()
return output_save_dir
View File
+267
View File
@@ -0,0 +1,267 @@
# Released under MIT license
# Copyright (c) 2022 finetuneanon (NovelAI/Anlatan LLC)
import numpy as np
import pickle
import time
def get_prng(seed):
return np.random.RandomState(seed)
class BucketManager:
def __init__(self, aspect_ratios, valid_ids=None, max_size=(768,512), divisible=64, step_size=8, min_dim=256, base_res=(512,512), bsz=1, world_size=1, global_rank=0, max_ar_error=4, seed=42, dim_limit=2048, debug=False):
self.res_map = aspect_ratios
if valid_ids is not None:
new_res_map = {}
valid_ids = set(valid_ids)
for k, v in self.res_map.items():
if k in valid_ids:
new_res_map[k] = v
self.res_map = new_res_map
self.max_size = max_size
self.f = 8
self.max_tokens = (max_size[0]/self.f) * (max_size[1]/self.f)
self.div = divisible
self.min_dim = min_dim
self.dim_limit = dim_limit
self.base_res = base_res
self.bsz = bsz
self.world_size = world_size
self.global_rank = global_rank
self.max_ar_error = max_ar_error
self.prng = get_prng(seed)
epoch_seed = self.prng.tomaxint() % (2**32-1)
self.epoch_prng = get_prng(epoch_seed) # separate prng for sharding use for increased thread resilience
self.epoch = None
self.left_over = None
self.batch_total = None
self.batch_delivered = None
self.debug = debug
self.gen_buckets()
self.assign_buckets()
self.start_epoch()
def gen_buckets(self):
if self.debug:
timer = time.perf_counter()
resolutions = []
aspects = []
w = self.min_dim
while (w/self.f) * (self.min_dim/self.f) <= self.max_tokens and w <= self.dim_limit:
h = self.min_dim
got_base = False
while (w/self.f) * ((h+self.div)/self.f) <= self.max_tokens and (h+self.div) <= self.dim_limit:
if w == self.base_res[0] and h == self.base_res[1]:
got_base = True
h += self.div
if (w != self.base_res[0] or h != self.base_res[1]) and got_base:
resolutions.append(self.base_res)
aspects.append(1)
resolutions.append((w, h))
aspects.append(float(w)/float(h))
w += self.div
h = self.min_dim
while (h/self.f) * (self.min_dim/self.f) <= self.max_tokens and h <= self.dim_limit:
w = self.min_dim
got_base = False
while (h/self.f) * ((w+self.div)/self.f) <= self.max_tokens and (w+self.div) <= self.dim_limit:
if w == self.base_res[0] and h == self.base_res[1]:
got_base = True
w += self.div
resolutions.append((w, h))
aspects.append(float(w)/float(h))
h += self.div
res_map = {}
for i, res in enumerate(resolutions):
res_map[res] = aspects[i]
self.resolutions = sorted(res_map.keys(), key=lambda x: x[0] * 4096 - x[1])
self.aspects = np.array(list(map(lambda x: res_map[x], self.resolutions)))
self.resolutions = np.array(self.resolutions)
if self.debug:
timer = time.perf_counter() - timer
print(f"resolutions:\n{self.resolutions}")
print(f"aspects:\n{self.aspects}")
print(f"gen_buckets: {timer:.5f}s")
def assign_buckets(self):
if self.debug:
timer = time.perf_counter()
self.buckets = {}
self.aspect_errors = []
skipped = 0
skip_list = []
for post_id in self.res_map.keys():
w, h = self.res_map[post_id]
aspect = float(w)/float(h)
bucket_id = np.abs(self.aspects - aspect).argmin()
if bucket_id not in self.buckets:
self.buckets[bucket_id] = []
error = abs(self.aspects[bucket_id] - aspect)
if error < self.max_ar_error:
self.buckets[bucket_id].append(post_id)
if self.debug:
self.aspect_errors.append(error)
else:
skipped += 1
skip_list.append(post_id)
for post_id in skip_list:
del self.res_map[post_id]
if self.debug:
timer = time.perf_counter() - timer
self.aspect_errors = np.array(self.aspect_errors)
print(f"skipped images: {skipped}")
print(f"aspect error: mean {self.aspect_errors.mean()}, median {np.median(self.aspect_errors)}, max {self.aspect_errors.max()}")
for bucket_id in reversed(sorted(self.buckets.keys(), key=lambda b: len(self.buckets[b]))):
print(f"bucket {bucket_id}: {self.resolutions[bucket_id]}, aspect {self.aspects[bucket_id]:.5f}, entries {len(self.buckets[bucket_id])}")
print(f"assign_buckets: {timer:.5f}s")
def start_epoch(self, world_size=None, global_rank=None):
if self.debug:
timer = time.perf_counter()
if world_size is not None:
self.world_size = world_size
if global_rank is not None:
self.global_rank = global_rank
# select ids for this epoch/rank
index = np.array(sorted(list(self.res_map.keys())))
index_len = index.shape[0]
index = self.epoch_prng.permutation(index)
index = index[:index_len - (index_len % (self.bsz * self.world_size))]
#print("perm", self.global_rank, index[0:16])
index = index[self.global_rank::self.world_size]
self.batch_total = index.shape[0] // self.bsz
assert(index.shape[0] % self.bsz == 0)
index = set(index)
self.epoch = {}
self.left_over = []
self.batch_delivered = 0
for bucket_id in sorted(self.buckets.keys()):
if len(self.buckets[bucket_id]) > 0:
self.epoch[bucket_id] = np.array([post_id for post_id in self.buckets[bucket_id] if post_id in index], dtype=np.int64)
self.prng.shuffle(self.epoch[bucket_id])
self.epoch[bucket_id] = list(self.epoch[bucket_id])
overhang = len(self.epoch[bucket_id]) % self.bsz
if overhang != 0:
self.left_over.extend(self.epoch[bucket_id][:overhang])
self.epoch[bucket_id] = self.epoch[bucket_id][overhang:]
if len(self.epoch[bucket_id]) == 0:
del self.epoch[bucket_id]
if self.debug:
timer = time.perf_counter() - timer
count = 0
for bucket_id in self.epoch.keys():
count += len(self.epoch[bucket_id])
print(f"correct item count: {count == len(index)} ({count} of {len(index)})")
print(f"start_epoch: {timer:.5f}s")
def get_batch(self):
if self.debug:
timer = time.perf_counter()
# check if no data left or no epoch initialized
if self.epoch is None or self.left_over is None or (len(self.left_over) == 0 and not bool(self.epoch)) or self.batch_total == self.batch_delivered:
self.start_epoch()
found_batch = False
batch_data = None
resolution = self.base_res
while not found_batch:
bucket_ids = list(self.epoch.keys())
if len(self.left_over) >= self.bsz:
bucket_probs = [len(self.left_over)] + [len(self.epoch[bucket_id]) for bucket_id in bucket_ids]
bucket_ids = [-1] + bucket_ids
else:
bucket_probs = [len(self.epoch[bucket_id]) for bucket_id in bucket_ids]
bucket_probs = np.array(bucket_probs, dtype=np.float32)
bucket_lens = bucket_probs
bucket_probs = bucket_probs / bucket_probs.sum()
bucket_ids = np.array(bucket_ids, dtype=np.int64)
if bool(self.epoch):
chosen_id = int(self.prng.choice(bucket_ids, 1, p=bucket_probs)[0])
else:
chosen_id = -1
if chosen_id == -1:
# using leftover images that couldn't make it into a bucketed batch and returning them for use with basic square image
self.prng.shuffle(self.left_over)
batch_data = self.left_over[:self.bsz]
self.left_over = self.left_over[self.bsz:]
found_batch = True
else:
if len(self.epoch[chosen_id]) >= self.bsz:
# return bucket batch and resolution
batch_data = self.epoch[chosen_id][:self.bsz]
self.epoch[chosen_id] = self.epoch[chosen_id][self.bsz:]
resolution = tuple(self.resolutions[chosen_id])
found_batch = True
if len(self.epoch[chosen_id]) == 0:
del self.epoch[chosen_id]
else:
# can't make a batch from this, not enough images. move them to leftovers and try again
self.left_over.extend(self.epoch[chosen_id])
del self.epoch[chosen_id]
assert(found_batch or len(self.left_over) >= self.bsz or bool(self.epoch))
if self.debug:
timer = time.perf_counter() - timer
print(f"bucket probs: " + ", ".join(map(lambda x: f"{x:.2f}", list(bucket_probs*100))))
print(f"chosen id: {chosen_id}")
print(f"batch data: {batch_data}")
print(f"resolution: {resolution}")
print(f"get_batch: {timer:.5f}s")
self.batch_delivered += 1
return (batch_data, resolution)
def generator(self):
if self.batch_delivered >= self.batch_total:
self.start_epoch()
while self.batch_delivered < self.batch_total:
yield self.get_batch()
if __name__ == "__main__":
# prepare a pickle with mapping of dataset IDs to resolutions called resolutions.pkl to use this
with open("resolutions.pkl", "rb") as fh:
ids = list(pickle.load(fh).keys())
counts = np.zeros((len(ids),)).astype(np.int64)
id_map = {}
for i, post_id in enumerate(ids):
id_map[post_id] = i
bm = BucketManager("resolutions.pkl", debug=True, bsz=8, world_size=8, global_rank=3)
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
print("got: " + str(bm.get_batch()))
bm = BucketManager("resolutions.pkl", bsz=8, world_size=1, global_rank=0, valid_ids=ids[0:16])
for _ in range(16):
bm.get_batch()
print("got from future epoch: " + str(bm.get_batch()))
bms = []
for rank in range(16):
bm = BucketManager("resolutions.pkl", bsz=8, world_size=16, global_rank=rank)
bms.append(bm)
for epoch in range(5):
print(f"epoch {epoch}")
for i, bm in enumerate(bms):
print(f"bm {i}")
first = True
for ids, res in bm.generator():
if first and i == 0:
#print(ids)
first = False
for post_id in ids:
counts[id_map[post_id]] += 1
print(np.bincount(counts))
+11 -67
View File
@@ -9,65 +9,6 @@ import signal
import time
import numpy as np
SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
SDXL_TURBO_MODEL_CACHE = "./models/SDXL_turbo.safetensors"
SDXL_TURBO_URL = "https://huggingface.co/stabilityai/sdxl-turbo/resolve/main/sd_xl_turbo_1.0_fp16.safetensors?download=true"
SDXL_LIGHTNING_MODEL_CACHE = "./models/SDXL_lightning.safetensors"
SDXL_LIGHTNING_URL = "https://huggingface.co/ByteDance/SDXL-Lightning/resolve/main/sdxl_lightning_8step.safetensors?download=true"
# Define model paths and URLs in a dictionary
MODEL_DICT = {
"sdxl": {
"path": SDXL_MODEL_CACHE,
"url": SDXL_URL,
"version": "sdxl"
},
"sd15": {
"path": SD15_MODEL_CACHE,
"url": SD15_URL,
"version": "sd15"
},
"sdxl_turbo": {
"path": SDXL_TURBO_MODEL_CACHE,
"url": SDXL_TURBO_URL,
"version": "sdxl_turbo"
},
"sdxl_lightning": {
"path": SDXL_LIGHTNING_MODEL_CACHE,
"url": SDXL_LIGHTNING_URL,
"version": "sdxl_lightning"
}
}
def download_weights(url, dest):
start = time.time()
print("downloading url: ", url)
print("downloading to: ", dest, '...')
# Make sure the destination directory exists
dest_dir = os.path.dirname(dest)
if not os.path.exists(dest_dir):
os.makedirs(dest_dir)
try:
subprocess.check_call(["wget", "-q", "-O", dest, url])
except subprocess.CalledProcessError as e:
print("Error occurred while downloading:")
print("Exit status:", e.returncode)
print("Output:", e.output)
except Exception as e:
print("An unexpected error occurred:", e)
print(f"Downloading {url} took {time.time() - start} seconds")
def clean_filename(filename):
allowed_chars = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-_"
return ''.join(c for c in filename if c in allowed_chars)
@@ -155,11 +96,11 @@ def merge_datasets(path_A, path_B, out_path, token_names):
def make_validation_img_grid(img_folder):
def make_validation_img_grid(img_folder, rows = 2):
"""
find all the .jpg imgs in img_folder (template = *.jpg)
if >=4 validation imgs, create a 2x2 grid of them
if >=4 validation imgs, create a rows x n grid of them
otherwise just return the first validation img
"""
@@ -171,18 +112,21 @@ def make_validation_img_grid(img_folder):
# If less than 4 validation images, return path of the first one
return os.path.join(img_folder, validation_imgs[0])
else:
# If >= 4 validation images, create 2x2 grid
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:4]]
# If >= 4 validation images, create 2xn grid
n_imgs = len(validation_imgs) // rows * rows
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:n_imgs]]
n_cols = int(n_imgs / rows)
# Assuming all images are the same size, get dimensions of first image
width, height = imgs[0].size
# Create an empty image with 2x2 grid size
grid_img = Image.new("RGB", (2 * width, 2 * height))
grid_img = Image.new("RGB", (n_cols * width, rows * height))
# Paste the images into the grid
for i in range(2):
for j in range(2):
for i in range(n_cols):
for j in range(rows):
grid_img.paste(imgs.pop(0), (i * width, j * height))
# Save the new image
@@ -438,7 +382,7 @@ def download_and_prep_training_data(data_location, data_dir):
print("Downloading training data...")
# we're assuming the data is prived as pipe seperated urls to .zip files
for url in str(data_location).split('|'):
download(url, data_dir)
download(url.strip(), data_dir)
# Loop over all files in the data directory:
for filename in os.listdir(data_dir):
+14
View File
@@ -0,0 +1,14 @@
import ujson
import os
def save_as_json(dictionary_or_list, filename: str):
with open(filename, "w") as fp:
ujson.dump(dictionary_or_list, fp, indent=4)
def load_json(filename: str):
assert os.path.exists(filename), f"Could not find json file: {filename}"
with open(filename) as json_file:
data = ujson.load(json_file)
return data
-23
View File
@@ -1,23 +0,0 @@
def get_avg_lr(optimizer):
# Calculate the weighted average effective learning rate
total_lr = 0
total_params = 0
for group in optimizer.param_groups:
d = group['d']
lr = group['lr']
bias_correction = 1 # Default value
if group['use_bias_correction']:
beta1, beta2 = group['betas']
k = group['k']
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
effective_lr = d * lr * bias_correction
# Count the number of parameters in this group
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
total_lr += effective_lr * num_params
total_params += num_params
if total_params == 0:
return 0.0
else: return total_lr / total_params
-174
View File
@@ -1,174 +0,0 @@
import os, json
import torch
from safetensors.torch import load_file
from typing import Dict
from peft import PeftModel
from ..dataset_and_utils import TokenEmbeddingsHandler
from safetensors.torch import save_file
from .string import replace_in_string
'''
from diffusers.utils import (
convert_all_state_dict_to_peft,
convert_state_dict_to_diffusers,
convert_unet_state_dict_to_peft
)
'''
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
return training_args["trigger_text"] + ", " + prompt
lora_name_encapsulated = "<" + lora_name + ">"
trigger_text = training_args["trigger_text"]
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK"
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
# Fix common mistakes
fix_replacements = {
r",,": ",",
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
r"\s\.": ".",
r"\s,": ","
}
prompt = replace_in_string(prompt, fix_replacements)
if verbose:
print('-------------------------')
print("Adjusted prompt for LORA:")
print(orig_prompt)
print('-- to:')
print(prompt)
print('-------------------------')
return prompt
def patch_pipe_with_lora(pipe, lora_path):
"""
update the pipe with the lora model and the token embeddings
"""
pipe.unet = PeftModel.from_pretrained(pipe.unet, lora_path)
pipe.unet.merge_adapter()
# Load the textual_inversion token embeddings into the pipeline:
try: #SDXL
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
except: #SD15
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
embeddings_path = [f for f in os.listdir(lora_path) if f.endswith("embeddings.safetensors")][0]
handler.load_embeddings(os.path.join(lora_path, embeddings_path))
return pipe
def unet_attn_processors_state_dict(unet) -> Dict[str, torch.tensor]:
"""
Returns:
a state dict containing just the attention processor parameters.
"""
attn_processors = unet.attn_processors
attn_processors_state_dict = {}
for attn_processor_key, attn_processor in attn_processors.items():
for parameter_key, parameter in attn_processor.state_dict().items():
attn_processors_state_dict[
f"{attn_processor_key}.{parameter_key}"
] = parameter
return attn_processors_state_dict
def save_lora(output_dir, global_step, unet, embedding_handler, token_dict, args_dict, is_lora, unet_lora_parameters, unet_param_to_optimize_names):
"""
Save the LORA model to output_dir, optionally with some example images
"""
print(f"Saving checkpoint at step.. {global_step}")
os.makedirs(output_dir, exist_ok=True)
if not is_lora:
lora_tensors = {
name: param
for name, param in unet.named_parameters()
if name in unet_param_to_optimize_names
}
save_file(lora_tensors, f"{output_dir}/unet.safetensors",)
elif len(unet_lora_parameters) > 0:
unet.save_pretrained(save_directory = output_dir)
try:
concept_name = args_dict["name"].lower()
except:
concept_name = "eden_concept_lora"
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
concept_name = concept_name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
embedding_handler.save_embeddings(f"{output_dir}/{concept_name}_embeddings.safetensors",)
with open(f"{output_dir}/special_params.json", "w") as f:
json.dump(token_dict, f)
with open(f"{output_dir}/training_args.json", "w") as f:
json.dump(args_dict, f, indent=4)
-13
View File
@@ -1,13 +0,0 @@
def print_trainable_parameters(model, name = ''):
trainable_params = 0
all_param = 0
for _, param in model.named_parameters():
all_param += param.numel()
if param.requires_grad:
trainable_params += param.numel()
line_delimiter = "#" * 70
print('\n', line_delimiter)
print(
f"Trainable {name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
)
print(line_delimiter, '\n')
-120
View File
@@ -1,120 +0,0 @@
import random
import json
import os
import gc
import torch
from ..dataset_and_utils import load_models
from .lora import patch_pipe_with_lora, prepare_prompt_for_lora
from ..val_prompts import val_prompts
from diffusers import EulerDiscreteScheduler
from PIL import Image
def make_validation_img_grid(img_folder):
"""
find all the .jpg imgs in img_folder (template = *.jpg)
if >=4 validation imgs, create a 2x2 grid of them
otherwise just return the first validation img
"""
# Find all validation images
validation_imgs = sorted([f for f in os.listdir(img_folder) if f.endswith(".jpg")])
if len(validation_imgs) < 4:
# If less than 4 validation images, return path of the first one
return os.path.join(img_folder, validation_imgs[0])
else:
# If >= 4 validation images, create 2x2 grid
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:4]]
# Assuming all images are the same size, get dimensions of first image
width, height = imgs[0].size
# Create an empty image with 2x2 grid size
grid_img = Image.new("RGB", (2 * width, 2 * height))
# Paste the images into the grid
for i in range(2):
for j in range(2):
grid_img.paste(imgs.pop(0), (i * width, j * height))
# Save the new image
grid_img_path = os.path.join(img_folder, "validation_grid.jpg")
grid_img.save(grid_img_path)
return grid_img_path
@torch.no_grad()
def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0"):
random.seed(seed)
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
concept_mode = training_args["concept_mode"]
if concept_mode == "style":
validation_prompts_raw = random.sample(val_prompts['style'], n_imgs)
validation_prompts_raw[0] = ''
elif concept_mode == "face":
validation_prompts_raw = random.sample(val_prompts['face'], n_imgs)
validation_prompts_raw[0] = '<concept>'
else:
validation_prompts_raw = random.sample(val_prompts['object'], n_imgs)
validation_prompts_raw[0] = '<concept>'
reload_entire_pipeline = False
if reload_entire_pipeline: # reload the entire pipeline from disk and load in the lora module
print(f"Reloading entire pipeline from disk..")
gc.collect()
torch.cuda.empty_cache()
(pipeline,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet) = load_models(pretrained_model, device, torch.float16)
pipeline = pipeline.to(device)
pipeline = patch_pipe_with_lora(pipeline, lora_path)
else:
print(f"Re-using training pipeline for inference, just swapping the scheduler..")
pipeline = training_pipeline
training_scheduler = pipeline.scheduler
pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config)
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path) for prompt in validation_prompts_raw]
generator = torch.Generator(device=device).manual_seed(0)
pipeline_args = {
"negative_prompt": "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft",
"num_inference_steps": n_steps,
"guidance_scale": 7,
"height": render_size[0],
"width": render_size[1],
}
if is_lora > 0:
cross_attention_kwargs = {"scale": lora_scale}
else:
cross_attention_kwargs = None
for i in range(n_imgs):
pipeline_args["prompt"] = validation_prompts[i]
print(f"Rendering validation img with prompt: {validation_prompts[i]}")
image = pipeline(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
image.save(os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"), format="JPEG", quality=95)
# create img_grid:
img_grid_path = make_validation_img_grid(lora_path)
if not reload_entire_pipeline: # restore the training scheduler
pipeline.scheduler = training_scheduler
return validation_prompts_raw
-24
View File
@@ -1,24 +0,0 @@
def compute_snr(noise_scheduler, timesteps):
"""
Computes SNR as per
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
"""
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = alphas_cumprod**0.5
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
# Expand the tensors.
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
# Compute SNR.
snr = (alpha / sigma) ** 2
return snr
-14
View File
@@ -1,14 +0,0 @@
import re
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
+280
View File
@@ -0,0 +1,280 @@
import os
from typing import Dict, List, Optional, Tuple
import random
import numpy as np
import pandas as pd
import gc
import PIL
import torch
import torch.utils.checkpoint
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
from PIL import Image
from safetensors import safe_open
from safetensors.torch import save_file
from torch.utils.data import Dataset
from transformers import AutoTokenizer, PretrainedConfig
import torch.nn.functional as F
import matplotlib.pyplot as plt
dtype_map = {
"fp16": torch.float16,
"bf16": torch.bfloat16,
"fp32": torch.float32
}
import re
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
def fix_prompt(prompt: str):
if not prompt:
return prompt
# Remove extra commas and spaces, and fix space before punctuation
prompt = re.sub(r"\s+", " ", prompt) # Replace multiple spaces with a single space
prompt = re.sub(r",,", ",", prompt) # Replace double commas with a single comma
prompt = re.sub(r"\s?,\s?", ", ", prompt) # Fix spaces around commas
prompt = re.sub(r"\s?\.\s?", ". ", prompt) # Fix spaces around periods
return prompt.strip() # Remove leading and trailing whitespace
def seed_everything(seed: int):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def zipdir(path, ziph, extension = '.py'):
# Zip the directory
for root, dirs, files in os.walk(path):
for file in files:
if file.endswith(extension):
ziph.write(os.path.join(root, file),
os.path.relpath(os.path.join(root, file),
os.path.join(path, '..')))
def pick_best_gpu_id():
try:
# pick the GPU with the most free memory:
gpu_ids = [i for i in range(torch.cuda.device_count())]
print(f"# of visible GPUs: {len(gpu_ids)}")
gpu_mem = []
for gpu_id in gpu_ids:
free_memory, tot_mem = torch.cuda.mem_get_info(device=gpu_id)
gpu_mem.append(free_memory)
print("GPU %d: %d MB free" %(gpu_id, free_memory / 1024 / 1024))
if len(gpu_ids) == 0:
# no GPUs available, use CPU:
os.environ["CUDA_VISIBLE_DEVICES"] = ""
return None
best_gpu_id = gpu_ids[np.argmax(gpu_mem)]
# set this to be the active GPU:
os.environ["CUDA_VISIBLE_DEVICES"] = str(best_gpu_id)
print("Using GPU %d" %best_gpu_id)
return best_gpu_id
except Exception as e:
print(f'Error picking best gpu: {e}')
print(f'Falling back to GPU 0')
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
return 0
import psutil
def print_system_info():
try:
# Print GPU memory information
gpu_ids = [i for i in range(torch.cuda.device_count())]
for gpu_id in gpu_ids:
free_memory, total_memory = torch.cuda.mem_get_info(device=gpu_id)
print(f"GPU {gpu_id}: {free_memory // 1024 // 1024} of {total_memory // 1024 // 1024} Mb free")
# Print disk space information
disk_usage = psutil.disk_usage('/')
total_disk = disk_usage.total // (1024 * 1024)
used_disk = disk_usage.used // (1024 * 1024)
percent_disk_used = disk_usage.percent
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
# Print RAM information
virtual_mem = psutil.virtual_memory()
total_ram = virtual_mem.total // (1024 * 1024)
current_ram = virtual_mem.used // (1024 * 1024)
percent_ram_used = virtual_mem.percent
print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
except Exception as e:
print(f'Error in gathering system info: {str(e)}')
return
def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1, max_val=1, ymax_f = 0.75, color = 'blue'):
try:
os.makedirs(checkpoint_dir, exist_ok=True)
# Flatten and concatenate all parameters into a single tensor
all_params = torch.cat([p.data.view(-1) for p in parameters])
# count number of parameters:
n_params = len(all_params)
if n_params == 0 or n_params > 1e9:
return
norm = torch.norm(all_params)
# Convert to CPU for plotting
all_params_cpu = all_params.cpu().float().numpy()
# Plot histogram
plt.figure()
plt.hist(all_params_cpu, bins=bins, density=False, color = color)
plt.ylim(0, ymax_f * len(all_params_cpu.flatten()))
plt.xlim(min_val, max_val)
plt.xlabel('Weight Value')
plt.ylabel('Count')
plt.title(f'{name} (std: {np.std(all_params_cpu):.5f}, norm: {norm:.3f}, step {step:03d})')
plt.savefig(f"{checkpoint_dir}/{name}_hist_{step:04d}.png")
plt.close()
except:
print(f'Error plotting {name} histogram')
def plot_curve(value_dict, xlabel, ylabel, title, save_path, log_scale = False, y_lims = None):
plt.figure()
for key in value_dict.keys():
values = value_dict[key]
plt.plot(range(len(values)), values, label=key)
if log_scale:
plt.yscale('log') # Set y-axis to log scale
plt.xlabel(xlabel)
plt.ylabel(ylabel)
if y_lims is not None:
plt.ylim(y_lims[0], y_lims[1])
plt.title(title)
plt.legend()
plt.savefig(save_path)
plt.close()
# plot the learning rates:
def plot_lrs(learning_rate_dict, save_path='learning_rates.png'):
plt.figure()
for key in learning_rate_dict.keys():
lrs = learning_rate_dict[key]
if len(lrs) == 0:
continue
plt.plot(range(len(lrs)), lrs, label=key)
plt.yscale('log') # Set y-axis to log scale
plt.ylim(1e-6, 3e-3)
plt.xlabel('Step')
plt.ylabel('Learning Rate')
plt.title('Learning Rate Curves')
plt.legend()
plt.savefig(save_path)
plt.close()
# plot the learning rates:
def plot_grad_norms(grad_norms, save_path='grad_norms.png'):
plt.figure()
plt.plot(range(len(grad_norms['unet'])), grad_norms['unet'], label='unet')
for i in range(2):
try:
plt.plot(range(len(grad_norms[f'text_encoder_{i}'])), grad_norms[f'text_encoder_{i}'], label=f'text_encoder_{i}')
except:
pass
plt.yscale('log') # Set y-axis to log scale
plt.ylim(1e-6, 100.0)
plt.xlabel('Step')
plt.ylabel('Grad Norm')
plt.title('Gradient Norms')
plt.legend()
plt.savefig(save_path)
plt.close()
def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dict = {}):
plt.figure()
anchor_values = []
for key in token_std_dict.keys():
tokenizer_i_token_stds = token_std_dict[key]
for i in range(len(tokenizer_i_token_stds)):
stds = tokenizer_i_token_stds[i]
if len(stds) == 0:
continue
anchor_values.append(stds[0])
encoder_index = int(key.split('_')[-1])
plt.plot(range(len(stds)), stds, label=f'{key}_tok_{i}', linestyle='dashed' if encoder_index > 0 else 'solid')
plt.xlabel('Step')
plt.ylabel('Token Embedding Std')
centre_value = 0.013
up_f, down_f = 1.4, 1.3
try:
plt.ylim(centre_value/down_f, centre_value*up_f)
except:
pass
# Plotting target values as horizontal lines
for label, value in target_value_dict.items():
plt.axhline(y=value, color='r', linestyle='-' if '0' in label else '--', label=label)
plt.text(0, value, label, ha='left', va='center')
plt.title('Token Embedding Std')
plt.legend()
plt.savefig(save_path)
plt.close()
from scipy.signal import savgol_filter
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'):
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red'}
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
plot_smoothed = ['img_loss']
plt.figure(figsize=(8, 5))
for key, losses in loss_dict.items():
if key == 'tot_loss':
continue
losses = np.array(losses)
if len(losses) < window_length:
continue
if key in plot_smoothed:
losses = savgol_filter(losses, window_length, polyorder)
label = f'Smoothed {key}'
linestyle = 'dashed'
else:
label = key
linestyle = 'solid'
plot_losses = losses / np.max(losses)
color = colormap.get(key, default_color) # Use the default color if the key is not in the colormap
plt.plot(plot_losses, label=label, color=color, linestyle=linestyle)
# Create the title:
title = 'Loss values:'
for key in values_to_add_to_title:
if key in loss_dict:
if loss_dict[key]:
title += f' {key}: {loss_dict[key][-1]:.3f}'
plt.title(title)
plt.xlabel('Optimizer Step')
plt.ylabel('Training Losses')
plt.ylim(0, 1.1) # Adjust the y-axis limits for normalized data
plt.legend(loc='lower left')
plt.savefig(save_path)
plt.close()
@@ -31,26 +31,21 @@ val_prompts['style'] = [
]
val_prompts["face"] = [
"an intricate wood carving of <concept> in a historic temple",
'<concept> as pixel art, 8-bit video game style',
'painting of <concept> by Vincent van Gogh',
'<concept> as a superhero, wearing a cape',
'<concept> as a statue made of marble',
'<concept> as a character in a noir graphic novel, under a rain-soaked streetlamp',
'stop motion animation of <concept> using clay, Wallace and Gromit style',
'<concept> portrayed in a famous renaissance painting, replacing Mona Lisas face',
'a photo of <concept> attending the Oscars, walking down the red carpet with sunglasses',
'<concept> as a pop vinyl figure, complete with oversized head and small body',
#'<concept> as a pop vinyl figure, complete with oversized head and small body',
'<concept> as a retro holographic sticker, shimmering in bright colors',
'<concept> as a bobblehead on a car dashboard, nodding incessantly',
"<concept> captured in a snow globe, complete with intricate details",
"a photo of <concept> climbing mount Everest in the snow, alpinism",
"<concept> as an action figure superhero, lego toy, toy story",
'a photo of a massive statue of <concept> in the middle of the city',
'a masterful oil painting portraying <concept> with vibrant colors, brushstrokes and textures',
'a vibrant low-poly artwork of <concept>, rendered in SVG, vector graphics',
'<concept>, polaroid photograph',
'a huge <concept> sand sculpture on a sunny beach, made of sand',
'an old, vintage, polaroid photograph of <concept>, artsy look',
'<concept> immortalized as an exquisite marble statue with masterful chiseling, swirling marble patterns and textures',
]
-70
View File
@@ -1,70 +0,0 @@
from trainer import TrainerConfig, Trainer
from preprocess import preprocess
import os
from io_utils import MODEL_DICT
out_root_dir = "./lora_models"
run_name = "face_01"
concept_mode = "face"
output_dir = os.path.join(out_root_dir, run_name)
input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
output_dir,
concept_mode = concept_mode,
input_zip_path = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
#caption_text="in the style of TOK, ",
caption_text="",
mask_target_prompts=None,
target_size=1024,
crop_based_on_salience=True,
use_face_detection_instead=False,
temp=0.7,
left_right_flip_augmentation=False,
augment_imgs_up_to_n = 20,
seed = 0,
caption_model = "blip"
)
print('-------------------------------------------')
print(f"Trigger text: {trigger_text}")
print(f'n_imgs: {n_imgs}')
print(f'concept_mode: {concept_mode}')
print('-------------------------------------------')
config = TrainerConfig(
pretrained_model = MODEL_DICT['sdxl'],
name='unnamed',
concept_mode=concept_mode,
trigger_text=trigger_text,
instance_data_dir = os.path.join(input_dir, "captions.csv"),
output_dir = output_dir,
resolution= 1024,
train_batch_size = 4,
max_train_steps = 600,
checkpointing_steps = 200,
num_train_epochs = 10000,
gradient_accumulation_steps = 1,
textual_inversion_lr = 5e-4,
textual_inversion_weight_decay = 3e-4,
lora_weight_decay = 0.00,
prodigy_d_coef = 1.0,
l1_penalty = 0.0,
snr_gamma = 5.0,
precision = "bf16",
token_dict = {"TOK": "<s0><s1>"},
inserting_list_tokens = ["<s0>","<s1>"],
is_lora = True,
lora_rank = 12,
lora_alpha = 12,
hard_pivot = False,
off_ratio_power = 0.1,
args_dict = {},
debug = True,
seed = 0
)
trainer = Trainer(config)
trainer.train()
print("DONE")
-27
View File
@@ -1,27 +0,0 @@
#!/bin/bash
# Check if the target directory is provided
if [ -z "$1" ]; then
echo "Usage: $0 target_directory"
exit 1
fi
# Check if the target directory exists
if [ ! -d "$1" ]; then
echo "Error: directory '$1' doesn't exist."
exit 1
fi
# Navigate to the target directory
cd "$1" || exit 1
# Initialize empty zip file
zip -r9 "args.zip" --exclude=*
# Find and zip all training_args.json files
find . -type f -name 'training_args.json' -exec zip -r "args.zip" {} +
# Navigate back to the original directory
cd - || exit 1
echo "All training_args.json files have been zipped into args.zip"