114 Commits
Author SHA1 Message Date
Xander Steenbrugge ca057f3597 Merge pull request #15 from ComfyNodePRs/licence-update
Update PyProject Toml - License
2025-08-04 23:16:38 +02:00
Xander Steenbrugge 9820c92bf3 Merge pull request #16 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-08-04 23:15:53 +02:00
snomiao c55440b95f chore(publish): update workflow permissions and action version for node publishing 2025-03-08 18:14:15 +00:00
snomiao d5f61bfd1d chore(licence-update): Update PyProject Toml - License 2025-03-08 17:24:58 +00:00
Xander Steenbrugge 686a36126b Merge pull request #13 from aurel-g/main
Require triton<3.2.0
2025-02-24 17:18:16 +01:00
aurel-g 59117c905f Require triton<3.2.0 2025-02-19 13:04:31 +01:00
xander e83a7971e5 bugfixes 2025-02-13 11:26:24 +01:00
xander 8fab8fe112 update preprocess bug 2025-02-12 23:59:02 +01:00
aiXander acb7c7bdc8 add training_args 2025-02-12 10:12:38 -08:00
Xander Steenbrugge f676f708cc Update requirements.txt 2024-12-20 11:46:20 +01:00
xander 4df66e932d tiny changes 2024-09-12 10:06:38 +02:00
xander c1916918a1 fix typo 2024-09-06 16:45:54 +02:00
xander 176971db00 add prompt_modifier 2024-09-03 17:24:18 +02:00
xander 1568f10e6e cleanup 2024-09-03 12:20:40 +02:00
xander 3e35201c87 update base_lr 2024-09-03 11:15:23 +02:00
xander f621670210 minor update 2024-09-03 10:20:34 +02:00
xander 2359aca6a9 update caching 2024-08-30 12:00:18 +02:00
xander cf8c0cd5e2 offload to disk when dataset is large 2024-08-29 15:07:50 +02:00
xander 413e4d898c fix img loading with transparency 2024-08-24 13:36:25 +02:00
xander 3b666a0b63 reset caption detail for florence 2024-08-24 13:29:12 +02:00
Xander Steenbrugge 77a10d5912 Update README.md 2024-08-23 11:17:58 +02:00
Xander Steenbrugge 6c3cf66d1f Update README.md 2024-08-23 11:17:27 +02:00
Xander Steenbrugge 98fc2d293b Update README.md 2024-08-23 11:13:04 +02:00
xander bde4dc1992 delete test 2024-08-23 11:11:48 +02:00
xander 92e3fe8599 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-08-23 11:05:53 +02:00
xander ae4f6a0ee1 add license 2024-08-23 11:05:45 +02:00
Xander Steenbrugge 622303c48d Update README.md 2024-08-23 11:02:56 +02:00
xander 7c0e1cbc62 update train workflow 2024-08-23 11:01:37 +02:00
xander e3f97bde50 update pyprompt.toml 2024-08-20 14:17:24 +02:00
xander cb03981dbe update name formatting 2024-08-20 13:31:54 +02:00
xander 6568a8f124 add toml 2024-08-19 20:14:14 +02:00
xander 09af837549 remove toml 2024-08-19 20:13:44 +02:00
xander 2be2e9b7b1 add pyproject.toml 2024-08-19 20:13:02 +02:00
Xander Steenbrugge d5998e8fe6 Merge pull request #1 from ComfyNodePRs/pyproject
Add pyproject.toml for Custom Node Registry
2024-08-19 20:12:44 +02:00
xander c396f9cd73 Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-08-19 20:10:16 +02:00
xander e642930cea save prompts to disk 2024-08-19 20:10:12 +02:00
Xander Steenbrugge 5798cbb187 Merge pull request #2 from ComfyNodePRs/publish
Add Github Action for Publishing to Comfy Registry
2024-08-19 19:59:57 +02:00
aiXander acb40ea1d3 update dockerignore 2024-08-19 03:26:46 -07:00
xander d263a95eb1 add name cleaning 2024-08-17 12:58:53 +02:00
snomiao 270e80364a chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-08-16 13:59:13 +00:00
snomiao 869896466a chore(publish): Add Github Action for Publishing to Comfy Registry 2024-08-16 13:59:12 +00:00
aiXander 2b7aff7b12 fix lora tar name 2024-08-16 02:40:30 -07:00
aiXander f3daa5bdeb ready for cog push 2024-08-15 04:49:24 -07:00
xander 414d43db2d update defaults 2024-08-15 13:39:07 +02:00
xander bf067c8b98 update sweep 2024-08-15 03:43:43 +02:00
aiXander 9c154e8804 add Eden_SDXL path 2024-08-14 18:32:53 -07:00
aiXander 4b44d89722 update print function 2024-08-14 18:10:26 -07:00
aiXander c82df696a3 update base lr 2024-08-14 18:06:10 -07:00
aiXander e04541805c remove line from config 2024-08-14 17:39:43 -07:00
aiXander 9920a0ea5a updates 2024-08-14 17:37:00 -07:00
aiXander 14cb9ac335 commit ti changes 2024-08-14 17:12:47 -07:00
aiXander 4ca8695d86 mini changes 2024-08-14 17:09:38 -07:00
aiXander 22ac54a66d update ti settings 2024-08-14 16:34:59 -07:00
aiXander 88946e88c2 more testing 2024-08-14 16:16:56 -07:00
aiXander 3d5bf750ed update msg 2024-08-14 15:39:28 -07:00
aiXander e1bd34af91 revert back to blip 2024-08-14 15:34:59 -07:00
aiXander 87aa2e6872 update reqs 2024-08-14 15:27:32 -07:00
aiXander 682edd9333 fix florence 2024-08-14 15:21:41 -07:00
xander 2434d4846d update loss 2024-08-15 00:00:40 +02:00
xander c89b22d034 update predict 2024-08-14 23:57:12 +02:00
xander 42e49afd78 update defaults 2024-08-14 23:52:29 +02:00
xander 5e4dcc2ece update workflows 2024-08-14 23:31:11 +02:00
xander 957420575e put in debug control 2024-08-14 22:16:24 +02:00
xander 3f7c620f5a tiny bugfixes 2024-08-14 22:15:08 +02:00
xander 715afb1f65 update note 2024-08-14 22:05:20 +02:00
xander 6b6d6a6ad1 update traininig workflow 2024-08-14 21:49:52 +02:00
xander 6eed725f94 update node settings 2024-08-14 21:48:29 +02:00
xander 6fa8dd47c1 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-08-14 21:32:21 +02:00
xander 2eda4dcba2 update node 2024-08-14 21:32:19 +02:00
xander 7627108bd6 update defaults 2024-08-14 21:32:07 +02:00
xander 9c63ec2d55 ready for new cog 2024-08-14 21:23:50 +02:00
xander f2c1a42254 update defaults 2024-08-14 20:04:14 +02:00
xander 32291a3b2f update defaults 2024-08-14 19:06:35 +02:00
xander 590c8577d7 update defaults 2024-08-14 12:19:21 +02:00
xander a83027d28c update defaults 2024-08-14 12:18:50 +02:00
xander 85facac79e update hyperparam sweep settings 2024-08-12 21:53:02 +02:00
xander ce9c608361 freeze bugfix 2024-08-12 16:01:47 +02:00
xander ff5cd50c49 add new params 2024-08-12 15:51:34 +02:00
xander e9555cd584 updates 2024-08-12 14:13:10 +02:00
aiXander 2e7879d39f update predict.py 2024-08-08 13:41:47 -07:00
Xander Steenbrugge c8d961dedf optionally set unet_optimizer to none 2024-08-08 18:57:20 +02:00
xander c0175fb67b update config 2024-08-07 23:44:57 +02:00
xander cb98a8e5ab update workflows 2024-08-07 23:36:18 +02:00
xander 2de4cbdd17 update workflow 2024-08-07 23:27:22 +02:00
xander 345615c6be remove training data zip 2024-08-07 23:26:06 +02:00
xander 7693899e8b rename imgs 2024-08-07 23:19:29 +02:00
xander a4edb04deb add dummy training data 2024-08-07 23:15:55 +02:00
xander d3d07bba0f update training workflow 2024-08-07 23:06:01 +02:00
xander 2b7e174efd update disable_ti option in node 2024-08-07 23:00:18 +02:00
xander c542e73f5a updates 2024-08-07 22:56:54 +02:00
xander c527fc690d update grid_img display 2024-08-07 22:21:45 +02:00
xander c31b122686 update readme 2024-08-07 22:14:03 +02:00
Xander Steenbrugge e7aa728177 Update README.md 2024-08-07 22:10:39 +02:00
Xander Steenbrugge 0a22b76511 Update README.md 2024-08-07 22:10:14 +02:00
Xander Steenbrugge 6e044fb167 Update README.md 2024-08-07 22:08:11 +02:00
Xander Steenbrugge 2ffb7ab1f5 Update README.md 2024-08-07 22:05:59 +02:00
Xander Steenbrugge 8d553441d7 Update README.md 2024-08-07 22:04:51 +02:00
Xander Steenbrugge c9428c15d5 Update README.md 2024-08-07 22:04:10 +02:00
xander 1eabc979a2 update training workflow with note 2024-08-07 21:58:45 +02:00
xander 39d70c1bda update defaults 2024-08-07 21:53:03 +02:00
xander 4959c87ee7 rename workflows 2024-08-07 21:50:57 +02:00
xander cc467f04b7 put comfyui workflows in separate folder 2024-08-07 21:07:40 +02:00
xander f5b2569646 simplify 2024-08-07 21:03:57 +02:00
xander 7f80ca88b3 simplify 2024-08-07 20:55:09 +02:00
xander 5a2283b4c0 update defaults 2024-08-05 13:44:51 +02:00
xander 2c1bb4a3f9 bugfix 2024-08-05 13:13:38 +02:00
xander ac2dcf787c remove objects 2024-08-05 04:07:18 +02:00
xander ff165f32db more cleanup and tweaking 2024-08-05 04:06:23 +02:00
xander 8e71457f07 massive refactor and attention regularization optimizations 2024-08-04 20:10:40 +02:00
xander ef61d5f0e4 update defaults 2024-08-03 04:01:47 +02:00
xander 47a0d37b45 tweak params 2024-08-03 03:49:00 +02:00
xander 32836aa82f massive upgrades to token attention reg, add florence captioning 2024-08-03 03:26:39 +02:00
xander 2b9f9c0f0a Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-08-02 18:46:49 +02:00
xander 479465b47a add token attention regularization 2024-08-02 18:46:42 +02:00
46 changed files with 2225 additions and 1306 deletions
+1 -8
View File
@@ -22,11 +22,4 @@ 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
lora_models/*
+26
View File
@@ -0,0 +1,26 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'edenartlab' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+1
View File
@@ -17,6 +17,7 @@ gridsearch*
aesthetic_score_best_model.pth
# experiment folders:
scripts/plots
conditioning_spaces/
training_args_x_*.json
xander_configs/
@@ -0,0 +1,782 @@
{
"last_node_id": 26,
"last_link_id": 49,
"nodes": [
{
"id": 15,
"type": "Note Plus (mtb)",
"pos": {
"0": 301,
"1": -120,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
519.1570279742359,
181.43570138536967
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"## How to trigger the LoRa:\nEden's trainer trains a token by default (if not disabled) which triggers the concept.\n\nFor \"object\" and \"face\" mode you just refer to your concept directly with the token, eg: \"a photo of embedding:MY\\_NAME\\_embedding\"\n\nFor \"style\" mode, you just prepend \"in the style of embedding:MY\\_NAME\\_embedding\" in the beginning of your prompt!",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 13,
"type": "Note Plus (mtb)",
"pos": {
"0": -506,
"1": -142,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
684.8613808966753,
209.8075855491436
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"## How to use:\n\n1. Find the training folder in ComfyUI/outputs\n2. Go into the checkpoints folder\n3. Pick the best checkpoint by looking at the validation_grid\n4. Copy the ***_embeddings.safetensors file to ComfyUI/models/embeddings\n5. Copy the ***_LoRa.safetensors file to ComfyUI/models/loras\n6. Hit Refresh in your ComfyUI\n7. Adjust this workflow to load both of those!\n8. Tweak the lora strength and embedding token strength to get the best results!\n\n\nHave fun! :)",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1162,
188
],
"size": {
"0": 140,
"1": 46
},
"flags": {},
"order": 13,
"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": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-451,
162
],
"size": [
429.578383119695,
98
],
"flags": {},
"order": 2,
"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": [
"zavychromaxl_v90.safetensors"
]
},
{
"id": 9,
"type": "SaveImage",
"pos": [
1327,
188
],
"size": {
"0": 453.85968017578125,
"1": 407.6841125488281
},
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 9
}
],
"properties": {
"Node name for S&R": "SaveImage"
},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
859,
26
],
"size": [
271.419023112996,
106
],
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
1024,
1024,
1
]
},
{
"id": 10,
"type": "LoraLoader",
"pos": [
31,
163
],
"size": {
"0": 337.44464111328125,
"1": 126
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 10
},
{
"name": "clip",
"type": "CLIP",
"link": 12
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
24
],
"shape": 3,
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
25,
26
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LoraLoader"
},
"widgets_values": [
"Eden_Token_LoRa_sdxl_LoRa.safetensors",
0.7000000000000001,
0.7000000000000001
]
},
{
"id": 3,
"type": "KSampler",
"pos": [
863,
186
],
"size": [
267.5352926614355,
262
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 24
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 40
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 44
},
{
"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": [
1093427037772868,
"randomize",
30,
8,
"euler_ancestral",
"normal",
1
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
402,
124
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 25
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
43
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"(in the style of embedding:Eden_Token_LoRa_sdxl_embeddings:1.0), \n\nwhat do i say to make me exist, oriental mythical beasts, in the golden danish age, in the history of television in the style of light violet and light red, serge najjar, playful and whimsical, associated press photo, afrofuturism-inspired, alasdair mclellan, electronic media"
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 26
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
45
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
""
]
},
{
"id": 22,
"type": "ControlNetApplyAdvanced",
"pos": [
672,
741
],
"size": {
"0": 315,
"1": 166
},
"flags": {},
"order": 11,
"mode": 4,
"inputs": [
{
"name": "positive",
"type": "CONDITIONING",
"link": 43
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 45
},
{
"name": "control_net",
"type": "CONTROL_NET",
"link": 38
},
{
"name": "image",
"type": "IMAGE",
"link": 47
}
],
"outputs": [
{
"name": "positive",
"type": "CONDITIONING",
"links": [
40
],
"shape": 3,
"slot_index": 0
},
{
"name": "negative",
"type": "CONDITIONING",
"links": [
44
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "ControlNetApplyAdvanced"
},
"widgets_values": [
1,
0,
1
]
},
{
"id": 23,
"type": "ControlNetLoader",
"pos": [
255,
733
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 4,
"mode": 4,
"outputs": [
{
"name": "CONTROL_NET",
"type": "CONTROL_NET",
"links": [
38
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ControlNetLoader"
},
"widgets_values": [
"SDXL/controlnet-canny-sdxl-1.0/diffusion_pytorch_model_V2.safetensors"
]
},
{
"id": 25,
"type": "AIO_Preprocessor",
"pos": [
258,
857
],
"size": {
"0": 315,
"1": 82
},
"flags": {},
"order": 7,
"mode": 4,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 48
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
47,
49
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "AIO_Preprocessor"
},
"widgets_values": [
"CannyEdgePreprocessor",
512
]
},
{
"id": 26,
"type": "PreviewImage",
"pos": [
679,
958
],
"size": [
307.39388136076934,
34.317686638573264
],
"flags": {},
"order": 10,
"mode": 4,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 49
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 24,
"type": "LoadImage",
"pos": [
-112,
857
],
"size": [
315,
314
],
"flags": {},
"order": 5,
"mode": 4,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
48
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"000050.jpg",
"image"
]
}
],
"links": [
[
2,
5,
0,
3,
3,
"LATENT"
],
[
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"
],
[
24,
10,
0,
3,
0,
"MODEL"
],
[
25,
10,
1,
6,
0,
"CLIP"
],
[
26,
10,
1,
7,
0,
"CLIP"
],
[
38,
23,
0,
22,
2,
"CONTROL_NET"
],
[
40,
22,
0,
3,
1,
"CONDITIONING"
],
[
43,
6,
0,
22,
0,
"CONDITIONING"
],
[
44,
22,
1,
3,
2,
"CONDITIONING"
],
[
45,
7,
0,
22,
1,
"CONDITIONING"
],
[
47,
25,
0,
22,
3,
"IMAGE"
],
[
48,
24,
0,
25,
0,
"IMAGE"
],
[
49,
25,
0,
26,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Optional Controlnet",
"bounding": [
-180,
645,
1205,
536
],
"color": "#3f789e",
"font_size": 24,
"locked": false
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.683013455365071,
"offset": {
"0": 493.3539751601979,
"1": 181.8365096487169
}
}
},
"version": 0.4
}
+312
View File
@@ -0,0 +1,312 @@
{
"last_node_id": 21,
"last_link_id": 35,
"nodes": [
{
"id": 4,
"type": "Display Any (rgthree)",
"pos": [
899,
263
],
"size": {
"0": 349.7635803222656,
"1": 88.69296264648438
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 34,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 3,
"type": "Display Any (rgthree)",
"pos": [
898,
150
],
"size": {
"0": 347.36749267578125,
"1": 87.71485137939453
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 33,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 5,
"type": "Display Any (rgthree)",
"pos": [
897,
373
],
"size": {
"0": 349.3501892089844,
"1": 98.81207275390625
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 35,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 2,
"type": "PreviewImage",
"pos": [
1307,
121
],
"size": {
"0": 557.6970825195312,
"1": 461.9648742675781
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 32
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 18,
"type": "Note Plus (mtb)",
"pos": {
"0": 411,
"1": -278,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": {
"0": 1447.654541015625,
"1": 329.1551513671875
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"\n## [Eden](https://www.eden.art/)\nThis LoRa trainer was made by the team behind https://www.eden.art/, led by https://x.com/xsteenbrugge\n\nIf you make awesome stuff w this trainer, give us a shout at:\nhttps://x.com/eden_art_ or\nhttps://www.instagram.com/eden.art____/\n\n---\n\n---\n\n## SD15 + SDXL:\nThis trainer works for both SD15 and SDXL models, but the default settings are primarily tuned for SDXL models. By selecting a specific ckpt_name, the trainer will automatically know if its an SDXL or SD15 model!\n\n---\n\n### A note on Embeddings:\nThis trainer optionally trains a textual inversion token into the LoRa, this is highly recommended when using SDXL models, but also means you have to load that token embedding when doing inference! See the example workflows in the repo.\n\n\nThe default settings work great for SDXL, SD15 usually need more training steps (eg 800) and sometimes benefits from disabling ti_training.\n\n---\n\n### A note on captioning:\nImages will get automatically captioned. It is recommended to put a .env file in the root of this custom node repo with your OpenAI API key, if found that will trigger a prompt_cleanup function that significantly improves results!\n",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 19,
"type": "Note Plus (mtb)",
"pos": {
"0": 48,
"1": 106,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
325.27622360123536,
609
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"\n## A note on settings:\n\nIf you disbale\\_ti (not recommended) you'll get a normal LoRa that does not use a token embedding, in that case you typically need to train for a bit longer and increase the unet_lr.\n\n---\n\n\"training_images\" can both be a path to a local folder or a url to a public .zip file of imgs which will get downloaded.\n\n---\n\nYou can provide custom captions by placing a filename.txt file for each filename.jpg in the training_images folder\n\n---\n\nI highly recommend to keep the training resolution at either 512 or 768.\n\n---\n\nn_tokens = 1 is currently broken, need to fix that.\n\n---\n\nAn embedding + LoRa checkpoint will get saved every **save\\_checkpoint\\_every\\_n\\_steps**. Based on the sample image grid you can then pick the best checkpoint to use in your workflows!\n\n---\n\nSetting debug=True will save a bunch of additional graphs and visualizations to track whats happening during training for advanced users.",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 21,
"type": "Eden_LoRa_trainer",
"pos": [
413,
124
],
"size": [
412.5926450093166,
591.4499899627035
],
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "sample_images",
"type": "IMAGE",
"links": [
32
],
"shape": 3,
"slot_index": 0
},
{
"name": "lora_path",
"type": "STRING",
"links": [
33
],
"shape": 3,
"slot_index": 1
},
{
"name": "embedding_path",
"type": "STRING",
"links": [
34
],
"shape": 3,
"slot_index": 2
},
{
"name": "final_msg",
"type": "STRING",
"links": [
35
],
"shape": 3,
"slot_index": 3
}
],
"properties": {
"Node name for S&R": "Eden_LoRa_trainer"
},
"widgets_values": [
"https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
"style",
"Eden_Token_LoRa",
"zavychromaxl_v90.safetensors",
512,
4,
300,
0.001,
0.0005,
16,
false,
3,
200,
6,
0.7,
false,
40092,
"randomize"
]
}
],
"links": [
[
32,
21,
0,
2,
0,
"IMAGE"
],
[
33,
21,
1,
3,
0,
"*"
],
[
34,
21,
2,
4,
0,
"*"
],
[
35,
21,
3,
5,
0,
"*"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.6830134553650705,
"offset": [
106.80403405562511,
337.0121827057539
]
}
},
"version": 0.4
}
+43
View File
@@ -0,0 +1,43 @@
Open Source Native License (OSNL)
Version 0.1 - March 1, 2024
Preamble
The Open Source Native License (OSNL) is designed to ensure that software remains free and open, fostering innovation and knowledge sharing within the community. It grants individuals, researchers, and commercial entities who open source their primary business assets, the freedom to use the software in any manner they choose.
This distinctive approach aims to balance the benefits of open-source development with the realities of commercial enterprise. It ensures that software remains a shared, community-driven resource while enabling businesses to thrive in an open-source ecosystem. Additional licenses are available for non-open source commercial entities.
1. Definitions
- "This License" refers to Version 1.0 of the Open Source Native License.
- "The Program" refers to the software distributed under this License.
- "You" refers to the individual or entity utilizing or contributing to the Program.
- "Primary Business Assets" are the core resources, capabilities, and technology that constitute the main value proposition and operational basis of your business.
2. Grant of License
Subject to the terms and conditions of this License, you are hereby granted a free, perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable license to use, reproduce, modify, distribute, and sublicense the Program, provided you comply with the following condition:
- Individual or researcher: You are granted the rights to use, modify, distribute, and contribute to the Program for any purpose, including educational, research, and personal projects, without the necessity to make your personal projects open source, provided these activities do not constitute a commercial enterprise. For any use that transitions to commercial purposes, the conditions applicable to commercial entities as outlined in this License will then apply.
- Commercial Entity who meets open source condition: Your primary business assets, including all core technologies, software, and platforms, must be available under an OSI-approved open source license or OSNL. This condition does not apply to ancillary or peripheral services not constituting primary business assets.
2.1 Commercial Use by Non-Open Source Businesses
Non-open source businesses that wish to utilize the Program or its derivatives as a component of their products or services are required to obtain an additional license. These entities must proactively contact Banodoco to request such a license. Banodoco reserves the right, at its own discretion, to grant or deny this additional license. Until an additional license is granted by Banodoco, non-open source businesses are not authorized to exercise any rights provided under this License regarding the use of the Program or its derivatives.
3. Redistribution
You may reproduce and distribute copies of the Program or derivative works thereof in any medium, with or without modifications, provided that you meet the following conditions:
- You must give any recipients of the Program a copy of this License.
- You must ensure that any modified files carry prominent notices stating that you changed the files.
- You must disclose the source of the Program, and if you distribute any portion of it in a compiled or object code form, you must also provide the full source code under this License.
- Any distribution of the Program or derivative works must comply with the Primary Business Open Source Condition.
4. Disclaimer of Warranty
THE PROGRAM IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES, OR OTHER LIABILITY ARISING FROM THE USE OF THE PROGRAM.
5. General
This License does not grant permission to use the trade names, trademarks, service marks, or product names of the Licensor, except as required for reasonable and customary use in describing the origin of the Program.
+26 -47
View File
@@ -1,10 +1,11 @@
# Trainer
This trainer was developed by the [**Eden** team](https://eden.art/)
This trainer was developed by the [**Eden** team](https://eden.art/), you can try our hosted version of the trainer in [**our app**](https://app.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.
The outputs of this trainer are fully compatible with ComfyUI and AUTO111, see documentation [here](https://docs.eden.art/docs/guides/concepts/#exporting-loras-for-use-in-other-tools).
A full guide on training can be found in [**our docs**](https://docs.eden.art/docs/guides/concepts/#training).
<p align="center">
<strong>Training images:</strong><br>
@@ -16,11 +17,26 @@ The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
</p>
The trainer supports 3 default modes:
### The trainer can be run in 4 different ways:
- [**as a hosted service on our website**](https://app.eden.art/)
- [**as a hosted service through replicate**](https://replicate.com/edenartlab/sdxl-lora-trainer)
- **as a ComfyUI node**
- **as a standalone python script**
### Using in ComfyUI:
- Example workflows for how to run the trainer and do inference with it can be found in `/ComfyUI_workflows`
- Importantly this trainer uses a chatgpt call to cleanup the auto-generated prompts and inject the trainable token, this will only work if you have a .env file containing your OPENAI key in the root of the repo dir that contains a single line: `OPENAI_API_KEY=your_key_string` Everything will work without this, but results will be better if you set this up, especially for 'face' and 'object' modes.
### 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.
<p align="center">
<strong>Style training example:</strong><br>
<img src="assets/style_training_example.jpg" alt="Image 1" style="width:80%;"/>
</p>
## Setup
Install all dependencies using
@@ -48,58 +64,21 @@ sudo chmod +x /usr/local/bin/cog
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
```
## Full unet finetuning
When running this trainer in native python, you can also perform full unet finetuning using something like (adjust to your needs)
`python main.py train_configs/full_finetuning_example.json`
## TODO's
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:
- 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)
- figure out why training is 3x slower through comfyui node versus just running main.py as a python job..?
- 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
- 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
- integrate Flux / SD3
- Add multi-concept training (multiple things represented by multiple tokens, trained into a single LoRa)
- add stronger token regularization (eg CelebBasis spanning basis)
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- implement prompt-aligned: https://prompt-aligned.github.io/
Tuning Experiments:
- try-out conditioning noise injection during training to increase robustness
- 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?
- offset noise
- AB test Dora vs Lora
Binary file not shown.

After

Width:  |  Height:  |  Size: 430 KiB

+6 -5
View File
@@ -3,10 +3,11 @@ 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 lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
-i concept_mode="style" \
-i sd_model_version="sdxl" \
-i max_train_steps="360" \
-i caption_model="blip" \
-i debug="False" \
-i max_train_steps="300" \
-i sample_imgs_lora_scale="0.7" \
-i n_sample_imgs="6" \
-i debug="True" \
-i seed="0"
-533
View File
@@ -1,533 +0,0 @@
{
"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
}
-28
View File
@@ -1,28 +0,0 @@
import numpy as np
from PIL import Image
import os
def generate_random_color_image(width, height):
"""Generate an image of random color."""
color = np.random.randint(0, 256, (3,), dtype=np.uint8)
image = np.full((height, width, 3), color, dtype=np.uint8)
return Image.fromarray(image)
def save_images(num_images, width, height, directory):
"""Save a specified number of random color images."""
for i in range(num_images):
image = generate_random_color_image(width, height)
image.save(f"{directory}/random_color_image_{i+1}.png")
# Parameters
num_images = 40
width = 1024
height = 1024
directory = "random_images"
os.makedirs(directory, exist_ok=True)
# Generate and save images
save_images(num_images, width, height, directory)
print(f'Saved {num_images} random color images to {os.path.abspath(directory)}')
+85 -173
View File
@@ -1,4 +1,3 @@
import fnmatch
import math
import os
import time
@@ -12,19 +11,17 @@ 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.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer, compute_token_attention_loss
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.ti_cross_attn_loss import init_daam_loss, plot_token_attention_loss
from trainer.optimizer import (
OptimizerCollection,
@@ -38,6 +35,7 @@ def train(config: TrainingConfig):
seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
@@ -49,14 +47,28 @@ def train(config: TrainingConfig):
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
if not config.sample_imgs_lora_scale:
if config.sd_model_version == "sdxl":
config.sample_imgs_lora_scale = 0.75
else:
config.sample_imgs_lora_scale = 0.85
if not config.validation_img_size:
if config.sd_model_version == "sdxl":
config.validation_img_size = 1024
else:
config.validation_img_size = 768
print("xxxxxxxxxxxxxxxxxxx")
print(config.prompt_modifier)
config, input_dir = preprocess(
config,
working_directory=config.output_dir,
@@ -151,15 +163,18 @@ def train(config: TrainingConfig):
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
)
if config.unet_lr > 0.0:
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
)
else:
optimizer_unet = None
print_trainable_parameters(unet, model_name = 'unet')
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
@@ -170,31 +185,26 @@ def train(config: TrainingConfig):
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:
print("Final training captions:")
print(train_dataset.captions[:40])
# offload the vae to cpu and release memory:
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_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)
config.num_train_epochs = int(math.ceil(config.max_train_steps / len(train_dataloader)))
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
print(f"--- Num samples = {len(train_dataset)}")
@@ -217,22 +227,18 @@ def train(config: TrainingConfig):
# 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': []}
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': [], 'token_attention_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 values for cold (starting) optimizer lr:
base_unet_lr = 2.0e-4 if (config.is_lora and config.disable_ti) else 5.0e-5
# 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
if not config.is_lora:
base_unet_lr = 1.0e-5
#######################################################################################################
"""
@@ -246,7 +252,8 @@ def train(config: TrainingConfig):
)
optimizers = optimizer_collection.optimizers
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
if config.debug:
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:
@@ -259,15 +266,12 @@ def train(config: TrainingConfig):
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 config.ti_optimizer != "prodigy" and optimizers['textual_inversion'] is not None:
# Apply the exponential learning rate
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 1.7
# Apply freezing condition
if completion_f > config.freeze_ti_after_completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] = 0.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
@@ -279,16 +283,26 @@ def train(config: TrainingConfig):
if optimizers['unet'] is not None:
# Calculate the exponential factor
exp_factor = (config.unet_lr / base_lr) ** (global_step / config.unet_lr_warmup_steps)
exp_factor = (config.unet_lr / base_unet_lr) ** (global_step / config.unet_lr_warmup_steps)
# Apply the exponential learning rate
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
optimizers['unet'].param_groups[0]['lr'] = base_unet_lr * exp_factor
if completion_f < config.freeze_unet_before_completion_f:
optimizers['unet'].param_groups[0]['lr'] = 0.0
if not config.aspect_ratio_bucketing:
captions, vae_latent, mask = batch
else:
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
mask = mask.to(config.device)
captions = list(captions)
if config.caption_dropout > 0.0:
for i in range(len(captions)):
if np.random.rand() < config.caption_dropout:
captions[i] = config.token_dict["TOK"]
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
config, pipe, captions
)
@@ -321,112 +335,15 @@ def train(config: TrainingConfig):
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 not config.disable_ti:
token_attention_loss = compute_token_attention_loss(pipe, embedding_handler, captions, mask, daam_loss)
losses['token_attention_loss'].append(token_attention_loss.item())
loss = loss + config.token_attention_loss_w * token_attention_loss
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:
@@ -442,7 +359,6 @@ def train(config: TrainingConfig):
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()
@@ -476,9 +392,14 @@ def train(config: TrainingConfig):
for std_i, std in enumerate(embedding_stds):
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
if global_step % 50 == 0 and not config.disable_ti and config.debug:
img_ratio = config.train_img_size[0] / config.train_img_size[1]
plot_token_attention_loss(config.output_dir, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio)
# Print some statistics:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)): #and global_step > 0:
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
config.save_as_json(
@@ -499,18 +420,6 @@ def train(config: TrainingConfig):
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)
@@ -519,8 +428,8 @@ def train(config: TrainingConfig):
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')
#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,
@@ -530,6 +439,8 @@ def train(config: TrainingConfig):
is_lora = config.is_lora,
pretrained_model = config.pretrained_model,
lora_scale = config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs,
device = config.device,
checkpoint_folder = None
@@ -543,10 +454,9 @@ def train(config: TrainingConfig):
images_done += config.train_batch_size
global_step += 1
if global_step % (config.max_train_steps//50) == 0:
if global_step % (config.max_train_steps//100) == 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:
@@ -611,6 +521,8 @@ def train(config: TrainingConfig):
is_lora=config.is_lora,
pretrained_model=config.pretrained_model,
lora_scale=config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs,
n_steps = 30,
device = config.device,
@@ -653,4 +565,4 @@ if __name__ == "__main__":
for progress in train(config=config):
print(f"Progress: {(100*progress):.2f}%", end="\r")
print("Training done :)")
print("Training done :)")
+52 -34
View File
@@ -19,19 +19,21 @@ class Eden_LoRa_trainer:
return {
"required": {
"training_images_folder_path": ("STRING", {"default": "."}),
"mode": (["style", "face", "object"], {"default": "style"}),
"lora_name": ("STRING", {"default": "Eden_Token_LoRa"}),
"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}),
"training_resolution": ("INT", {"default": 512, "min": 256, "max": 1024}),
"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}),
"max_train_steps": ("INT", {"default": 300, "min": 10, "max": 10000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0, "max": 0.005, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.0005, "min": 0.0, "max": 0.005, "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}),
"disable_ti": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 3, "min": 1, "max": 5}),
"save_checkpoint_every_n_steps": ("INT", {"default": 200, "min": 10, "max": 10000}),
"n_sample_imgs": ("INT", {"default": 4, "min": 2, "max": 10}),
"sample_imgs_lora_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.25}),
"plot_training_graphs_on_disk": ("BOOLEAN", {"default": False}),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
}
}
@@ -44,25 +46,28 @@ class Eden_LoRa_trainer:
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,
lora_name,
mode,
training_resolution,
train_batch_size,
max_train_steps ,
ti_lr,
unet_lr,
lora_rank,
disable_ti,
n_tokens,
plot_training_graphs_on_disk,
save_checkpoint_every_n_steps,
n_sample_imgs,
sample_imgs_lora_scale,
seed,
):
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("FLORENCE", os.path.join(folder_paths.models_dir, "LLM"))
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"))
@@ -71,22 +76,26 @@ class Eden_LoRa_trainer:
config = TrainingConfig(
name=lora_name,
output_dir="output",
lora_training_urls=training_images_folder_path,
concept_mode=mode,
ckpt_path=ckpt_path,
seed=seed,
resolution=resolution,
resolution=training_resolution,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=checkpointing_steps,
checkpointing_steps=save_checkpoint_every_n_steps,
n_sample_imgs=(n_sample_imgs//2) * 2,
sample_imgs_lora_scale=sample_imgs_lora_scale,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
use_dora=use_dora,
use_dora=False,
caption_model="blip",
disable_ti=disable_ti,
n_tokens=n_tokens,
verbose=True,
debug=debug_mode,
debug=plot_training_graphs_on_disk,
)
pbar = comfy.utils.ProgressBar(100)
@@ -101,8 +110,6 @@ class Eden_LoRa_trainer:
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
@@ -120,11 +127,22 @@ class Eden_LoRa_trainer:
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,]
# Load the grid images:
grid_images = []
grid_dir = os.path.dirname(output_save_dir)
for f in os.listdir(grid_dir):
if "validation_grid" in f:
grid_image = Image.open(os.path.join(grid_dir, f))
grid_image = np.array(grid_image).astype(np.float32) / 255.0
grid_image = torch.from_numpy(grid_image)
grid_images.append(grid_image)
grid_images = torch.stack(grid_images)
# Make sure that grid_images always has 4 dimensions:
if len(grid_images.shape) == 3:
grid_images = grid_images.unsqueeze(0)
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)
return (grid_images, lora_path, embedding_path, final_msg)
+35 -18
View File
@@ -14,7 +14,6 @@ from main import train
from typing import Iterator, Optional
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
@@ -66,19 +65,19 @@ class Predictor(BasePredictor):
),
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
default=300
),
checkpointing_steps: int = Input(
description="Save a checkpoint every n steps (The final checkpoint will always be saved)",
default=10000
),
resolution: int = Input(
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 768",
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
default=0.0003
),
ti_lr: float = Input(
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
@@ -88,13 +87,25 @@ class Predictor(BasePredictor):
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
ge=1, le=4, default=3
),
train_batch_size: int = Input(
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
default=4
),
n_sample_imgs: int = Input(
description="Number of sample images in validation grid",
default=4
),
validation_img_size: int = Input(
description="Resolution of sample images in validation grid",
default=1024
),
sample_imgs_lora_scale: float = Input(
description="Scale factor for LoRa when generating sample images. If not provided, will be set automatically",
default=None
),
seed: int = Input(
description="Random seed for reproducible training. Leave empty to use a random seed",
@@ -124,13 +135,15 @@ class Predictor(BasePredictor):
sd_model_version=sd_model_version,
seed=seed,
resolution=resolution,
validation_img_size=[validation_img_size, validation_img_size],
sample_imgs_lora_scale=sample_imgs_lora_scale,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=10000,
checkpointing_steps=checkpointing_steps,
n_sample_imgs=n_sample_imgs,
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,
@@ -162,9 +175,13 @@ 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")
comfy_workflows_path = "ComfyUI_workflows"
if os.path.exists(comfy_workflows_path) and os.path.isdir(comfy_workflows_path):
for root, dirs, files in os.walk(comfy_workflows_path):
for file in files:
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, os.path.dirname(comfy_workflows_path))
tar.add(file_path, arcname=arcname)
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
+16
View File
@@ -0,0 +1,16 @@
[project]
name = "Eden LoRa Trainer"
description = "Highly optimized and fast LoRa trainer for SD15 and SDXL models. Does LoRa training, textual inversion and full finetuning!"
version = "1.0.0"
license = { text = "Other" }
dependencies = ["torch==2.1.0", "torchaudio==2.1.0", "torchvision==0.16.0", "transformers==4.38.0", "diffusers==0.29.2", "tokenizers==0.15.2", "huggingface-hub==0.23.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", "torchtyping==0.1.5", "einops==0.8.0", "timm==1.0.8"]
[project.urls]
Repository = "https://github.com/edenartlab/sd-lora-trainer"
Website = "https://www.eden.art/"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "eden"
DisplayName = "Eden LoRa Trainer"
Icon = ""
+11 -7
View File
@@ -1,10 +1,10 @@
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
torch>=2.1.0
torchaudio>=2.1.0
torchvision>=0.16.0
transformers>=4.38.0
diffusers>=0.29.2
tokenizers>=0.15.2
huggingface-hub==0.23.2
ujson==5.10.0
scipy==1.14.0
peft==0.10.0
@@ -20,3 +20,7 @@ omegaconf==2.3.0
ujson==5.10.0
bitsandbytes==0.43.1
setuptools==70.3.0
torchtyping==0.1.5
einops==0.8.0
timm==1.0.8
triton<3.2.0
+31 -37
View File
@@ -1,10 +1,8 @@
"""
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
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip
Objects:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
@@ -14,7 +12,12 @@ https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets
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
https://edenartlab-lfs.s3.amazonaws.com/datasets/clipx.zip
/home/rednax/Documents/datasets/good_styles/eden_crystals
/home/rednax/Documents/datasets/good_styles/does2
/home/rednax/Documents/datasets/good_styles/beeple
"""
@@ -35,12 +38,12 @@ def hamming_distance(dict1, dict2):
#######################################################################################
# Setup the base experiment config:
exp_name = "beeple"
exp_name = "ygor_sd15"
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
n_exp = 100 # how many random experiment settings to generate
min_hamming_distance = 2 # min_n_params that have to be different from any previous experiment to be scheduled
nohup = False
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values
@@ -48,45 +51,39 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sdxl"],
"sd_model_version": ["sd15"],
"lora_training_urls": [
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
"/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean"
],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.8],
"disable_ti": ['false', 'true'],
"sample_imgs_lora_scale": [0.9],
"caption_dropout": [0.2],
"seed": [0],
"resolution": [512],
"train_batch_size": [4],
"resolution": [512,640,768],
"train_batch_size": [8],
"n_sample_imgs": [8],
"max_train_steps": [1200],
"checkpointing_steps": [200],
"max_train_steps": [2000],
"checkpointing_steps": [500],
"gradient_accumulation_steps": [1],
"n_tokens": [2],
"ti_lr": [0.001],
"ti_weight_decay": [0.001],
"l1_penalty": [0.0],
"n_tokens": [3],
"disable_ti": ['true'],
"ti_lr": [0.0001],
"token_warmup_steps": [0],
"tok_cov_reg_w": [2000],
"unet_lr": [0.001, 0.0003],
"lora_rank": [8,24,64],
"use_dora": ['false', 'true'],
"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'],
"unet_optimizer_type": ['adamw'],
"is_lora": ['true'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["blip", "gpt4-v"],
"caption_model": ["florence", "blip", "no_caption"],
"augment_imgs_up_to_n": [40],
"verbose": ['true'],
"debug": ['true']
@@ -103,7 +100,7 @@ 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
try_sampling_n_times = 120
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
resamples, combination = 0, None
@@ -125,9 +122,6 @@ for exp_index in tqdm(range(n_exp)): # number of combinations you want to gener
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
+199
View File
@@ -0,0 +1,199 @@
import os
import json
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from collections import defaultdict
from sklearn.metrics import r2_score
# Step 2: Define a function to count JPG files in a directory
def count_jpg_files(directory):
return len([f for f in os.listdir(directory) if f.lower().endswith('.jpg')])
# Step 3: Define a function to find and load the training_args.json file
def load_training_args(directory):
for root, dirs, files in os.walk(directory):
if 'training_args.json' in files:
with open(os.path.join(root, 'training_args.json'), 'r') as f:
return json.load(f)
return None
# Step 4: Traverse the directory structure and collect data
def collect_data(root_dir, mode):
data = []
for root, dirs, files in os.walk(root_dir):
if "checkpoints" in dirs:
checkpoints_dir = os.path.join(root,"checkpoints")
if mode == "final_checkpoint":
checkpoint_dirs = sorted([f for f in os.listdir(checkpoints_dir) if os.path.isdir(os.path.join(checkpoints_dir, f))])
checkpoint_dir = os.path.join(checkpoints_dir, checkpoint_dirs[-1])
score = count_jpg_files(checkpoint_dir)
training_args = load_training_args(checkpoint_dir)
if mode == "n_validation_grids":
score = count_jpg_files(checkpoints_dir)
training_args = load_training_args(checkpoints_dir)
if not training_args:
print(f"Warning: Could not load training_args.json from {checkpoints_dir}")
continue
data.append((training_args, score))
print(f"Collected data from {len(data)} runs")
return data
# Step 5: Process the collected data to identify varying hyperparameters
def identify_varying_hyperparams(data, skip_params=['output_dir', 'start_time', 'name']):
all_params = set().union(*[set(args.keys()) for args, _ in data])
varying_params = {}
def make_hashable(val):
if isinstance(val, dict):
return tuple(sorted((k, make_hashable(v)) for k, v in val.items()))
elif isinstance(val, list):
return tuple(make_hashable(v) for v in val)
elif isinstance(val, set):
return frozenset(make_hashable(v) for v in val)
return val
for param in all_params:
if param in skip_params:
continue
try:
values = [make_hashable(args.get(param)) for args, _ in data if param in args]
unique_values = set(values)
if len(unique_values) > 1:
# Check if all values are numeric
try:
numeric_values = [float(v) for v in unique_values]
varying_params[param] = set(numeric_values)
except ValueError:
# If not all numeric, keep as is
varying_params[param] = unique_values
print(f"---> Parameter '{param}' varies across runs")
# Special handling for dictionary-type parameters
if all(isinstance(v, dict) for v in values):
print(f"Dictionary values for '{param}':")
for v in unique_values:
print(f" {v}")
elif len(unique_values) <= 5: # Print up to 5 unique values
print(f"Unique values: {unique_values}")
else:
print(f"Number of unique values: {len(unique_values)}")
except TypeError as e:
print(f"Warning: Could not process values for parameter '{param}'. Error: {e}")
#print(f"Values: {[args.get(param) for args, _ in data if param in args]}")
return varying_params
def create_plots(data, varying_params, outdir, top = 0.15):
os.makedirs(outdir, exist_ok=True)
for param, values in varying_params.items():
if all(isinstance(v, dict) for v in values):
print(f"Skipping plot for dictionary parameter '{param}'")
continue
plt.figure(figsize=(12, 8))
param_data = defaultdict(list)
all_scores = []
for args, score in data:
if param in args:
value = args[param]
value_str = str(value)
param_data[value_str].append(score)
all_scores.append(score)
# Calculate global top threshold
global_top_percent = np.percentile(all_scores, 100*(1-top))
top_75_percent = np.percentile(all_scores, 25)
# Sort the values
try:
values_list = sorted(param_data.keys(), key=float)
except ValueError:
values_list = sorted(param_data.keys())
all_x = []
all_y = []
for i, value_str in enumerate(values_list):
scores = param_data[value_str]
# Apply jitter
jittered_x = np.random.normal(i, 0.1, size=len(scores))
jittered_y = np.array(scores) + np.random.normal(0, 0.01 * max(scores), size=len(scores))
# Use global top 25% threshold
top_mask = np.array(scores) >= global_top_percent
# Emphasize top 25% scores using JITTERED coordinates for plotting
sns.scatterplot(x=jittered_x[top_mask], y=jittered_y[top_mask], alpha=0.6, color='black', marker='X', s=40, linewidth=1)
# Plot all scores using JITTERED coordinates
sns.scatterplot(x=jittered_x, y=jittered_y, alpha=0.6, label=value_str)
all_x.extend([i] * len(scores))
all_y.extend(scores)
# Calculate trendline for top 75% data
x = np.array(all_x)
y = np.array(all_y)
top_75_mask = y >= top_75_percent
x_75 = x[top_75_mask]
y_75 = y[top_75_mask]
z = np.polyfit(x_75, y_75, 1)
p = np.poly1d(z)
plt.plot(range(len(values_list)), p(range(len(values_list))), "r--", alpha=0.8,
label=f'Top 75%: y={z[0]:.2f}x+{z[1]:.2f}\nR²: {r2_score(y_75, p(x_75)):.4f}')
# Calculate trendline for global top 25% scoring datapoints
top_mask = y >= global_top_percent
x_top = x[top_mask]
y_top = y[top_mask]
z_top = np.polyfit(x_top, y_top, 1)
p_top = np.poly1d(z_top)
plt.plot(range(len(values_list)), p_top(range(len(values_list))), "g--", alpha=0.8,
label=f'Top 25%: y={z_top[0]:.2f}x+{z_top[1]:.2f}\nR²: {r2_score(y_top, p_top(x_top)):.4f}')
plt.xlabel(param)
plt.ylabel('Score')
plt.title(f'Effect of {param} on Score')
# Set x-ticks and labels
plt.xticks(range(len(values_list)), values_list, rotation=45, ha='right')
# Adjust legend
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
plt.tight_layout()
# Save figure with error handling
try:
plt.savefig(f'{outdir}/{param}_vs_score.png', dpi=200, bbox_inches='tight')
except ValueError:
print(f"Warning: Failed to save image for {param}. Skipping...")
plt.close()
print(f"Plots have been saved as PNG files in {outdir}")
if __name__ == "__main__":
root_dir = "/home/rednax/SSD2TB/Github_repos/Eden/sd-lora-trainer/lora_models/ygor_sd15"
outdir = os.path.join('.', os.path.basename(root_dir))
# Collect data
#data = collect_data(root_dir, mode = "final_checkpoint")
data = collect_data(root_dir, mode = "n_validation_grids")
# Identify varying hyperparameters
varying_params = identify_varying_hyperparams(data)
# Create plots
create_plots(data, varying_params, outdir)
-8
View File
@@ -1,8 +0,0 @@
# 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
@@ -0,0 +1,25 @@
{
"name": "stitchly",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/stitchly_florence",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.9,
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 5000,
"checkpointing_steps": 1000,
"unet_optimizer_type": "AdamW8bit",
"disable_ti": true,
"is_lora": false,
"caption_dropout": 0.2,
"caption_model": "florence",
"ti_lr": 0.0001,
"unet_lr": 0.0001,
"lora_rank": 4,
"debug": true
}
-27
View File
@@ -1,27 +0,0 @@
{
"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
}
+14 -22
View File
@@ -1,28 +1,20 @@
{
"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,
"freeze_ti_after_completion_f": 1.0,
"name": "mira_sdxl_ti",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.7,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 900,
"checkpointing_steps": 300,
"caption_model": "blip",
"debug": true
}
+21
View File
@@ -0,0 +1,21 @@
{
"name": "banny",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip",
"concept_mode": "object",
"sample_imgs_lora_scale": 0.75,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 300,
"checkpointing_steps": 200,
"disable_ti": false,
"caption_model": "florence",
"ti_lr": 0.001,
"unet_lr": 0.0003,
"lora_rank": 16,
"debug": true
}
+10 -16
View File
@@ -1,27 +1,21 @@
{
"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",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 512,
"resolution": 768,
"train_batch_size": 4,
"n_sample_imgs": 6,
"n_sample_imgs": 8,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 300,
"checkpointing_steps": 200,
"disable_ti": false,
"caption_model": "florence",
"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,
"unet_lr": 0.0005,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+5 -16
View File
@@ -1,27 +1,16 @@
{
"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",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
"concept_mode": "face",
"seed": 1,
"sample_imgs_lora_scale": 0.7,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"n_sample_imgs": 8,
"max_train_steps": 300,
"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
}
+12 -19
View File
@@ -1,28 +1,21 @@
{
"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,
"name": "banny",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip",
"concept_mode": "object",
"sample_imgs_lora_scale": 0.75,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 800,
"token_warmup_steps": 0,
"n_sample_imgs": 8,
"max_train_steps": 300,
"checkpointing_steps": 200,
"disable_ti": false,
"caption_model": "florence",
"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,
"unet_lr": 0.0003,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+11 -17
View File
@@ -1,27 +1,21 @@
{
"name": "clipx_sd15",
"name": "twisting_realities_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 512,
"resolution": 768,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"n_sample_imgs": 8,
"max_train_steps": 1200,
"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,
"disable_ti": false,
"caption_model": "blip",
"ti_lr": 0.001,
"unet_lr": 0.0003,
"lora_rank": 16,
"debug": true
}
@@ -0,0 +1,21 @@
{
"name": "ygor_painting_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 0,
"resolution": 640,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 4000,
"checkpointing_steps": 500,
"disable_ti": true,
"caption_model": "blip",
"ti_lr": 0.001,
"unet_lr": 0.0015,
"lora_rank": 64,
"debug": true
}
+10 -17
View File
@@ -1,28 +1,21 @@
{
"name": "clipx_sdxl",
"name": "twisting_realities",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"sample_imgs_lora_scale": 0.75,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"max_train_steps": 300,
"checkpointing_steps": 200,
"disable_ti": false,
"caption_model": "florence",
"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,
"unet_lr": 0.0003,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+16
View File
@@ -0,0 +1,16 @@
{
"name": "xander_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/beeple",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 300,
"checkpointing_steps": 200,
"caption_model": "florence",
"debug": true
}
+16
View File
@@ -0,0 +1,16 @@
{
"name": "mira_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/01_people/mira",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 360,
"checkpointing_steps": 240,
"caption_model": "florence",
"debug": true
}
+25 -4
View File
@@ -54,10 +54,31 @@ def set_adapter_scales(pipe, lora_scale = 1.0):
return pipe
import re
def remove_delimiter_characters(name: str, max_length: int = 255) -> str:
# Define a regular expression pattern to match all weird and special characters
pattern = r'[^\w.-]+' # Matches any character that is not alphanumeric, underscore, dot, or hyphen
# Replace all occurrences of the pattern with a single underscore
cleaned_name = re.sub(pattern, '_', name)
# Replace multiple consecutive underscores with a single underscore
cleaned_name = re.sub(r'_+', '_', cleaned_name)
# Strip leading or trailing underscores and dots
cleaned_name = cleaned_name.strip('_.')
# Ensure the name doesn't start with a dot (to avoid hidden files on Unix)
cleaned_name = cleaned_name.lstrip('.')
# Truncate to max_length if necessary
cleaned_name = cleaned_name[:max_length]
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("|", "_")
# Raise an error if the name is empty or malformed after cleaning
if not cleaned_name:
raise ValueError("Malformed name")
return cleaned_name
# Convert to WebUI format
def convert_pytorch_lora_safetensors_to_webui(
@@ -184,7 +205,7 @@ def save_checkpoint(
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")
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_lora.safetensors")
)
else:
# Save the entire, finetuned unet weights:
+30 -35
View File
@@ -4,11 +4,13 @@ from pydantic import BaseModel
import json, time, os
from typing import Literal
from trainer.utils.utils import pick_best_gpu_id
from trainer.checkpoint import remove_delimiter_characters
class ModelPaths:
def __init__(self):
self.paths = {
"BLIP": "./cache",
"FLORENCE": "./cache",
"CLIP": "./cache",
"SR": "./cache",
"SD": "./models",
@@ -23,9 +25,9 @@ class ModelPaths:
model_paths = ModelPaths()
# Default download urls in case no local model is found:
# Default SD model download urls in case no local model is found:
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/Eden_SDXL.safetensors"
#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 = {
@@ -37,7 +39,9 @@ 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"
prompt_modifier: str = None # optional prompt modifier
caption_model: Literal["gpt4-v", "blip", "florence", "no_caption"] = "florence"
caption_dropout: float = 0.1 # dropout rate for captions: occasionally use empty prompt
sd_model_version: Literal["sdxl", "sd15", None] = None
ckpt_path: str = None # optional hardcoded checkpoint path
pretrained_model: dict = None
@@ -47,36 +51,36 @@ class TrainingConfig(BaseModel):
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
max_train_steps: int = 300
num_train_epochs: int = None
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
unet_lr: float = 0.0003
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
lora_weight_decay: float = 0.004
ti_lr: float = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
ti_lr: float = 0.001
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
freeze_ti_after_completion_f: float = 0.7 # freeze the TI after this fraction of the training is done
freeze_unet_before_completion_f: float = 0.0 # freeze the UNET before this fraction of the training is done
token_attention_loss_w: float = 3e-7
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
tok_cov_reg_w: float = 0. # regularizes the token covariance matrix wrt pretrained, normal tokens
l1_penalty: float = 0.03 # 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
lora_rank: int = 16
use_dora: bool = False
left_right_flip_augmentation: bool = True
@@ -91,17 +95,12 @@ class TrainingConfig(BaseModel):
debug: bool = False
allow_tf32: bool = True
disable_ti: bool = False
skip_gpt_cleanup: 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>"}
n_tokens: int = 3
inserting_list_tokens: List[str] = ["<s0>","<s1>","<s2>"]
token_dict: dict = {"TOK": "<s0><s1><s2>"}
device: str = "cuda:0"
crops_coords_top_left_h: int = 0
crops_coords_top_left_w: int = 0
do_cache: bool = True
unet_learning_rate: float = 1.0
lr_num_cycles: int = 1
lr_power: float = 1.0
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0
training_attributes: dict = {}
@@ -126,15 +125,14 @@ class TrainingConfig(BaseModel):
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.name = os.path.basename(self.lora_training_urls)[:40]
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}"
self.name = remove_delimiter_characters(self.name)
timestamp = datetime.now().strftime("%d%b_%H%M")
self.output_dir = self.output_dir + f"/{self.name}_{timestamp}-{self.concept_mode}_res{self.resolution}_{self.max_train_steps}steps"
os.makedirs(self.output_dir, exist_ok=True)
if self.seed is None:
@@ -143,17 +141,14 @@ class TrainingConfig(BaseModel):
if self.unet_lr_warmup_steps is None:
self.unet_lr_warmup_steps = self.max_train_steps
if self.checkpointing_steps < 1:
self.checkpointing_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.")
+35 -31
View File
@@ -1,6 +1,7 @@
import os
import torch
import numpy as np
from tqdm import tqdm
import pandas as pd
import PIL
from PIL import Image
@@ -32,7 +33,6 @@ class PreprocessedDataset(Dataset):
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,
@@ -42,13 +42,14 @@ class PreprocessedDataset(Dataset):
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)
self.data = pd.read_csv(self.csv_path, dtype={"caption": str})
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)
self.captions = self.captions.fillna("")
self.image_path = self.data["image_path"]
if "mask_path" not in self.data.columns:
@@ -62,23 +63,31 @@ class PreprocessedDataset(Dataset):
self.text_dropout = text_dropout
self.size = size
if do_cache:
print("Caching latents, masks and captions...\n")
# If the training data is small we can keep everything in memory, otherwise offload to disk
self.do_cache = True if len(self.data) < 500 else False
if self.do_cache:
print("Encoding latents, masks and captions and storing in memory...\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)
for idx in tqdm(range(len(self.data))):
vae_latent, mask, _ = self._process(idx)
self.vae_latents.append(vae_latent)
self.masks.append(mask)
self.masks.append(mask.detach())
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
del self.vae_encoder
else:
self.do_cache = False
else: # Store the latents and masks on disk
print("Encoding latents, masks and captions and storing on disk...\n")
self.vae_latents = None
self.masks = None
for idx in tqdm(range(len(self.data))):
vae_latent, mask, image_path = self._process(idx)
torch.save(vae_latent, os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
torch.save(mask, os.path.join(self.data_dir, f"{idx}_mask.pt"))
del self.vae_encoder
torch.cuda.empty_cache()
if aspect_ratio_bucketing:
print("Using aspect ratio bucketing.")
@@ -100,12 +109,9 @@ class PreprocessedDataset(Dataset):
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:
@@ -141,28 +147,24 @@ class PreprocessedDataset(Dataset):
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
dtype=self.vae_encoder.dtype
)
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
dtype=self.vae_encoder.dtype
)
vae_latent = self.vae_encoder.encode(image).latent_dist
vae_latent = self.vae_encoder.encode(image.to(self.vae_encoder.device)).latent_dist
dummy_vae_latent = vae_latent.sample()
if self.mask_path is None:
mask = torch.ones_like(
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
mask = torch.ones_like(dummy_vae_latent, dtype=self.vae_encoder.dtype)
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 = prepare_mask(mask, self.size[0], self.size[1]).to(dtype=self.vae_encoder.dtype)
mask_dtype = mask.dtype
mask = mask.float()
@@ -174,7 +176,7 @@ class PreprocessedDataset(Dataset):
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
return vae_latent, mask.squeeze()
return vae_latent, mask.squeeze(), image_path
def __getitem__(
self, idx: int, bucketing_resolution:tuple = None
@@ -182,10 +184,12 @@ class PreprocessedDataset(Dataset):
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)
return self.captions[idx], vae_latent.squeeze().detach(), self.masks[idx].detach()
else: # Load from disk:
vae_latent = torch.load(os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
vae_latent = vae_latent.sample() * self.vae_scaling_factor
return caption, vae_latent.squeeze(), mask
mask = torch.load(os.path.join(self.data_dir, f"{idx}_mask.pt"))
caption = self.captions[idx]
return caption, vae_latent.squeeze().detach(), mask.detach()
+1 -41
View File
@@ -260,11 +260,7 @@ class TokenEmbeddingsHandler:
# 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,
)
crops_coords_top_left = (0,0)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
@@ -397,7 +393,6 @@ class TokenEmbeddingsHandler:
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:
@@ -430,41 +425,6 @@ class TokenEmbeddingsHandler:
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])]
+10 -6
View File
@@ -155,11 +155,7 @@ def get_conditioning_signals(config, pipe, captions):
# 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,
)
crops_coords_top_left = (0,0)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
@@ -245,7 +241,7 @@ def encode_prompt_advanced(
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:
if lora_path and token_scale != 0:
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
else:
lora_prompt = prompt
@@ -299,6 +295,8 @@ def render_images(
is_lora,
pretrained_model,
lora_scale,
disable_ti=False,
prompt_modifier=None,
n_steps=25,
n_imgs=4,
device="cuda:0",
@@ -328,6 +326,9 @@ def render_images(
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
validation_prompts_raw[0] = "<concept>"
if prompt_modifier:
validation_prompts_raw = [prompt_modifier.format(prompt) for prompt in validation_prompts_raw]
if (
checkpoint_folder is not None
): # reload the entire pipeline from disk and load in the lora module
@@ -376,6 +377,7 @@ def render_images(
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
)
pipeline_args["prompt_embeds"] = c
@@ -415,6 +417,7 @@ def render_images_eval(
pretrained_model: dict,
trigger_text: str,
lora_scale=0.7,
disable_ti=False,
n_steps=25,
n_imgs=4,
device="cuda:0",
@@ -469,6 +472,7 @@ def render_images_eval(
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
)
pipeline_args["prompt_embeds"] = c
+76 -1
View File
@@ -4,6 +4,81 @@ 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
import torch.nn.functional as F
def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_loss, verbose=0):
"""
Custom loss function to regularize the attention maps of the token embeddings.
"""
masks = masks[:, 0].float()
img_ratio = masks.shape[-1] / masks.shape[-2]
att_L2_losses = []
ti_heatmaps = []
ti_masks = []
att_reg_threshold = 0.0
# attention_maps.shape = [n_layers, batch_size, w, h, 77]
attention_maps = daam_loss.process_and_stack_attention_scores(img_ratio)
n_layers, batch_size, w, h, n_tokens = attention_maps.shape
# reshape masks to match attention maps:
# masks.shape = [batch_size, w2, h2]
masks = F.interpolate(masks.unsqueeze(1), size=(attention_maps.shape[-3], attention_maps.shape[-2])).squeeze(1)
masks = masks.unsqueeze(0).unsqueeze(-1)
masks = masks.repeat(n_layers, 1, 1, 1, n_tokens)
for batch_index, caption in enumerate(captions):
token_indices = pipe.tokenizer.encode(caption)
# Penalize the mean attention score of each token:
mean_att_per_token = attention_maps[:,batch_index, :, :, 1:len(token_indices)-1].mean(dim=[0,1,2])
att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean()
att_L2_losses.append(att_L2_loss)
try:
ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids]
except:
continue
batch_ti_heatmaps, batch_ti_masks = [], []
# Extract the attention heatmaps corresponding to the trainable token embeddings:
for text_token_index in ti_token_indices:
ti_heatmap = attention_maps[:,batch_index, :, :, text_token_index].mean(dim=0)
ti_mask = masks[:,batch_index, :, :, text_token_index].mean(dim=0)
batch_ti_heatmaps.append(ti_heatmap.float())
batch_ti_masks.append(ti_mask)
ti_heatmaps.append(torch.stack(batch_ti_heatmaps))
ti_masks.append(torch.stack(batch_ti_masks))
if len(ti_heatmaps) == 0:
return torch.tensor(0.0).to(masks.dtype)
ti_heatmaps = torch.stack(ti_heatmaps)
ti_masks = torch.stack(ti_masks)
#ti_heatmaps.shape = [batch_size, n_tokens, w, h]
token_means = ti_heatmaps.mean(dim=[2,3])
token_attention_scores = token_means.var(dim=1)
# Avoid large attention scores in general:
reg_loss_0 = 5.0 * torch.stack(att_L2_losses).mean()
# Avoid large attention scores for ti tokens, inside the masked region:
reg_loss_1 = 1.0 * (torch.relu(ti_heatmaps * ti_masks)**2).mean()
# Avoid large attention scores for ti tokens, outside of the masked region:
reg_loss_2 = 2.0 * (torch.relu(ti_heatmaps * (1 - ti_masks) + 10)**2).mean()
# Make the Ti tokens have similar avg attention scores (equal distribution of concept information over tokens):
reg_loss_3 = 1.0 * token_attention_scores.mean()
if verbose:
print(f"reg_loss_0: {reg_loss_0.item():.4f}")
print(f"reg_loss_1: {reg_loss_1.item():.4f}")
print(f"reg_loss_2: {reg_loss_2.item():.4f}")
print(f"reg_loss_3: {reg_loss_3.item():.4f}")
return (reg_loss_0 + reg_loss_1 + reg_loss_2 + reg_loss_3).to(masks.dtype)
def compute_snr(noise_scheduler, timesteps):
"""
@@ -118,7 +193,7 @@ class ConditioningRegularizer:
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):
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.01, 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
+23 -19
View File
@@ -4,24 +4,30 @@ 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):
def load_models(pretrained_model, device, weight_dtype = torch.float16):
# check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
tokenizer_two, text_encoder_two = None, None
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
try:
print("Loading as SDXL model...")
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
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)
except:
print("Loading as SD15 model...")
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)
@@ -31,24 +37,11 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
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..?")
vae.to(device, dtype=weight_dtype)
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,
@@ -82,16 +75,27 @@ def download_weights(url, dest):
print(f"Downloading {url} took {time.time() - start} seconds")
def print_trainable_parameters(model, model_name = ''):
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()
def format_param_count(count):
if count < 1000:
return f"{count}"
elif count < 1_000_000:
return f"{count/1000:.1f}K"
else:
return f"{count/1_000_000:.1f}M"
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}%"
f"Trainable {model_name} params: {format_param_count(trainable_params)} "
f"|| All params: {format_param_count(all_param)} "
f"|| trainable = {100 * trainable_params / all_param:.2f}%"
)
print(line_delimiter)
+95 -81
View File
@@ -53,7 +53,7 @@ except:
# 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
MAX_GPT_PROMPTS = 80
def _find_files(pattern, dir="."):
"""Return list of files matching pattern in a given directory, in absolute format.
@@ -231,7 +231,6 @@ def clipseg_mask_generator(
return masks
import textwrap
def cleanup_prompts_with_chatgpt(
prompts,
@@ -332,12 +331,12 @@ def extract_gpt_concept_description(gpt_completion, concept_mode):
return concept_name
def post_process_captions(captions, text, concept_mode, job_seed):
def post_process_captions(captions, text, concept_mode, job_seed, skip_gpt_cleanup=False):
text = text.strip()
gpt_cleanup_worked = False
gpt_concept_description = None
if len(captions) >= MIN_GPT_PROMPTS 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) and not skip_gpt_cleanup:
retry_count = 0
while retry_count < 5:
try:
@@ -419,6 +418,7 @@ 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)
model.to("cpu")
del model
gc.collect()
torch.cuda.empty_cache()
@@ -440,51 +440,6 @@ 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,
@@ -495,6 +450,7 @@ def gpt4_v_caption_dataset(
return captions
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."
#prompt = "Concisely describe the main subject(s) in the image without assumptions with at most 20 words. Ignore the background and only describe the people / animals / objects in the foreground. Dont start with statements like 'The image features...', just describe what you see."
headers = {
"Content-Type": "application/json",
@@ -543,17 +499,84 @@ def gpt4_v_caption_dataset(
device = "cuda:0" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
@torch.no_grad()
def florence_caption_dataset(images, captions):
#workaround for unnecessary flash_attn requirement
from unittest.mock import patch
from transformers.dynamic_module_utils import get_imports
from transformers import AutoProcessor, AutoModelForCausalLM
def fixed_get_imports(filename: str | os.PathLike) -> list[str]:
if not str(filename).endswith("modeling_florence2.py"):
return get_imports(filename)
imports = get_imports(filename)
try:
imports.remove("flash_attn")
except:
pass
return imports
with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement
model = AutoModelForCausalLM.from_pretrained("microsoft/Florence-2-large", attn_implementation="sdpa", device_map=device, torch_dtype=torch_dtype,trust_remote_code=True)
processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large", trust_remote_code=True, cache_dir = model_paths.get_path("FLORENCE"))
for i, image in enumerate(tqdm(images)):
if captions[i] is None:
#prompt = random.choice(["<CAPTION>", "<DETAILED_CAPTION>"])
prompt = "<CAPTION>"
prompt = "<DETAILED_CAPTION>"
prompt = "<MORE_DETAILED_CAPTION>"
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
generated_ids = model.generate(
input_ids=inputs["input_ids"],
pixel_values=inputs["pixel_values"],
max_new_tokens=1024,
num_beams=random.choice([2,3,4])
)
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
parsed_answer = processor.post_process_generation(generated_text, task=prompt, image_size=(image.width, image.height))
caption = parsed_answer[prompt]
captions[i] = caption.replace("The image shows a ", "A ")
model.to('cpu')
del model
del processor
gc.collect()
torch.cuda.empty_cache()
return captions
@torch.no_grad()
def caption_dataset(
images: List[Image.Image],
captions: List[str],
caption_model: Literal[str] = "blip"
caption_model: Literal["blip", "gpt4-v", "florence"] = "blip"
) -> List[str]:
# if all captions are already generated, we dont need to do anything:
if all(captions):
print(f"All captions loaded from disk, skipping captioning...")
return captions
if "blip" in caption_model:
captions = blip_caption_dataset(images, captions)
elif "gpt4-v" in caption_model:
captions = gpt4_v_caption_dataset(images, captions)
elif "florence" in caption_model:
captions = florence_caption_dataset(images, captions)
else:
print("WARNING: not using any captions!")
captions = [""] * len(images)
gc.collect()
torch.cuda.empty_cache()
return captions
@@ -624,15 +647,15 @@ def random_crop(image, scale=(0.85, 0.95)):
return image.crop((left, top, left + new_width, top + new_height))
def gaussian_blur(image):
return image.filter(ImageFilter.GaussianBlur(radius=1))
def gaussian_blur(image, radius = 1.0):
return image.filter(ImageFilter.GaussianBlur(radius=radius))
def augment_image(image):
image = hue_augmentation(image)
image = color_jitter(image)
image = random_crop(image)
if random.random() < 0.5:
image = gaussian_blur(image)
image = gaussian_blur(image, radius = random.uniform(0.0, 1.0))
return image
def round_to_nearest_multiple(x, multiple):
@@ -720,9 +743,10 @@ def load_and_save_masks_and_captions(
n_length = len(files)
files = sorted(files)[:n_length]
images, captions = [], []
images, captions, img_paths = [], [], []
for file in files:
images.append(load_image_with_orientation(file))
img_paths.append(file)
caption_file = os.path.splitext(file)[0] + ".txt"
if os.path.exists(caption_file) and use_dataset_captions:
with open(caption_file, "r") as f:
@@ -763,44 +787,39 @@ def load_and_save_masks_and_captions(
upscale_margin = 0.75
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:
if add_lr_flips and len(images) < MAX_GPT_PROMPTS:
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)
print(f"Generating {len(images)} captions using {caption_model} in {concept_mode} mode...")
captions = caption_dataset(images, captions, caption_model = caption_model)
# Save captions back to disk:
for i, img_path in enumerate(img_paths):
caption_path = os.path.splitext(img_path)[0] + ".txt"
with open(caption_path, "w") as f:
f.write(captions[i])
# 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]
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 = [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)
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed, skip_gpt_cleanup=config.skip_gpt_cleanup)
if config.prompt_modifier:
print(config.prompt_modifier)
captions = [config.prompt_modifier.format(caption) for caption in captions]
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
# if we still have a 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])
@@ -809,19 +828,18 @@ def load_and_save_masks_and_captions(
images.extend(aug_imgs)
captions.extend(aug_caps)
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 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:
@@ -889,10 +907,6 @@ def load_and_save_masks_and_captions(
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
print("Final captions:")
for caption in captions:
print(caption)
# iterate through the images, masks, and captions and add a row to the dataframe for each
print("Saving final training dataset...")
for idx, (image, mask, caption) in enumerate(zip(images, seg_masks, captions)):
+128 -53
View File
@@ -2,15 +2,88 @@ from functools import reduce
from diffusers import StableDiffusionXLPipeline
from diffusers.models.attention_processor import AttnProcessor2_0, Attention
from typing import Optional
import os
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
import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.axes_grid1 import make_axes_locatable
def plot_token_attention_loss(folder, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio):
try:
batch_index = 0
timestep = timesteps[batch_index].item()
if timestep > 700:
batch_index += 1
timestep = timesteps[batch_index].item()
folder = os.path.join(folder, "attention_heatmaps")
os.makedirs(folder, exist_ok=True)
token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])]
plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens
# Calculate global min and max for consistent colormap
all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float()
for i in plot_token_indices]
vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps])
# Heatmap plots
fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices),
figsize=(3 * len(plot_token_indices), 10))
title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}"
fig.suptitle(title_str, fontsize=16)
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = all_heatmaps[idx]
im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax)
axes[idx].set_title(f"{token_strings[text_token_index]}")
axes[idx].axis("off")
# Add colorbar
divider = make_axes_locatable(axes[idx])
cax = divider.append_axes("bottom", size="5%", pad=0.05)
plt.colorbar(im, cax=cax, orientation="horizontal")
plt.tight_layout()
fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight')
plt.close(fig)
# Histogram plot
fig, ax = plt.subplots(figsize=(12, 6))
ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16)
ax.set_xlabel("Attention Value", fontsize=12)
ax.set_ylabel("Frequency", fontsize=12)
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = all_heatmaps[idx]
ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index],
alpha=0.5, density=True)
ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
ax.grid(alpha=0.3)
plt.tight_layout()
# Add token attention loss as text
plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}",
transform=ax.transAxes, ha='right', va='top',
bbox=dict(facecolor='white', edgecolor='black', alpha=0.8))
plt.ylim(0, 0.6)
fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight')
plt.close(fig)
except:
print("Failed to plot token attention loss")
# Find all instances of AttnProcessor2_0 in the UNet
def find_attnprocessor2_0(unet):
@@ -108,8 +181,6 @@ class DAAMLossAttnProcessor2_0:
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)
@@ -165,6 +236,37 @@ class DAAMLoss:
x.name for x in attention_processors
]
def process_and_stack_attention_scores(self, img_ratio: float):
reshaped_tensors = []
min_heatmap_pixels = np.inf
# Process each attention score
for processor in self.attention_processors:
score = processor.cross_attention_scores
bs, seq_len, channels = score.shape
# Calculate width and height based on img_ratio
width = round(math.sqrt(seq_len * img_ratio))
height = round(width / img_ratio)
reshaped_score = rearrange(score, 'b (h w) c -> b h w c', h=height, w=width)
reshaped_tensors.append(reshaped_score)
if height*width < min_heatmap_pixels:
min_heatmap_pixels = height*width
min_heatmap_shape = height, width
# Interpolate and standardize all tensors to the same size
for i, heatmap in enumerate(reshaped_tensors):
if heatmap.shape[1] * heatmap.shape[2] != min_heatmap_pixels:
# Interpolating to match the smallest tensor size uniformly
heatmap = F.interpolate(heatmap.permute(0, 3, 1, 2), size=(min_heatmap_shape[0], min_heatmap_shape[1]), mode='bicubic').permute(0, 2, 3, 1)
reshaped_tensors[i] = heatmap
# Stack all tensors along the first dimension
stacked_tensor = torch.stack(reshaped_tensors, dim=0)
return stacked_tensor
def get_all_cross_attention_scores(self):
cross_attention_scores = {}
@@ -175,76 +277,53 @@ class DAAMLoss:
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"]:
def get_image_heatmap(self, text_token_index: int, layer_name: str, img_ratio: float):
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
width = round(math.sqrt(cross_attention_scores_single_token.shape[-1] * img_ratio))
height = round(width / img_ratio)
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]))
height = height,
width = width
)
return heatmap
def get_the_daam_heatmap(self, text_token_index: int) ->TensorType["batch", "height", "width"]:
def get_the_daam_heatmap(self, text_token_index: int, img_ratio: float, resize = 'min'):
all_heatmaps = []
for layer_name in self.layer_names:
heatmap = self.get_image_heatmap(
text_token_index=text_token_index,
layer_name=layer_name
layer_name=layer_name,
img_ratio = img_ratio
)
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)
if resize == 'max':
heatmap_height = max(heatmap.shape[1] for heatmap in all_heatmaps)
heatmap_width = max(heatmap.shape[2] for heatmap in all_heatmaps)
elif resize == 'min':
heatmap_height = min(heatmap.shape[1] for heatmap in all_heatmaps)
heatmap_width = min(heatmap.shape[2] for heatmap in all_heatmaps)
## now resize all_heatmaps to (batch, max_height, max_width) using F.interpolate
## now resize all_heatmaps to (batch, heatmap_height, heatmap_width) using F.interpolate
resized_heatmaps = [
F.interpolate(input = x.unsqueeze(1), size = (max_height, max_width)).squeeze(1)
F.interpolate(input = x.unsqueeze(1), size = (heatmap_height, heatmap_width)).squeeze(1)
for x in all_heatmaps
]
return sum(resized_heatmaps)
resized_heatmaps = torch.stack(resized_heatmaps)
## Average the heatmaps:
avg_heatmap = resized_heatmaps.mean(dim=0)
return avg_heatmap
def get_module_by_name(module: nn.Module, name: str):
"""Retrieve a module nested in another by its access string."""
@@ -253,10 +332,8 @@ def get_module_by_name(module: nn.Module, name: str):
names = name.split(sep=".")
return reduce(getattr, names, module)
def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionXLPipeline, DAAMLoss]:
assert isinstance(pipeline, StableDiffusionXLPipeline)
def init_daam_loss(pipeline):
## find out where the attention processor thingies are
module_names = find_attnprocessor2_0(
unet = pipeline.unet
@@ -265,8 +342,6 @@ def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionX
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]
+12
View File
@@ -290,6 +290,18 @@ def load_image_with_orientation(path, mode="RGB"):
elif orientation == 8:
image = image.rotate(90, expand=True)
if image.mode == 'P':
image = image.convert('RGBA')
if image.mode == 'CMYK':
image = image.convert('RGB')
# Remove alpha channel if present
if image.mode in ('RGBA', 'LA'):
background = Image.new('RGB', image.size, (255, 255, 255))
background.paste(image, mask=image.split()[3]) # 3 is the alpha channel
image = background
# Convert to the desired mode
return image.convert(mode)
+1 -1
View File
@@ -237,7 +237,7 @@ def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dic
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'}
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red', 'token_attention_loss': 'purple'}
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
plot_smoothed = ['img_loss']
+2 -2
View File
@@ -2,9 +2,9 @@
val_prompts = {}
val_prompts['style'] = [
'a beautiful mountainous landscape, boulders, fresh water stream, setting sun',
'the stunning skyline of New York City',
'the stunning skyline of New York City, setting sun, skyscrapers, wallpaper',
'fruit hanging from a tree, highly detailed texture, soil, rain, drops, photo realistic, surrealism, highly detailed, 8k macrophotography',
'the Taj Mahal, stunning wallpaper',
'the Taj Mahal, stunning wallpaper, architecture, ancient, marble, white, intricate, detailed',
'A majestic tree rooted in circuits, leaves shimmering with data streams, stands as a beacon where the digital dawn caresses the fog-laden, binary soil—a symphony of pixels and chlorophyll.',
'A beautiful octopus, with swirling tendrils and a pulsating heart of fiery opal hues, hovers ethereally against a starry void, sculpted through a meticulous flame-working technique.',
'a stunning image of an aston martin sportscar',