149 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
mayukhdeb f21efd7840 train TI on 2 tokens 2024-07-29 10:11:15 -07:00
mayukhdeb 29855d94ba ignore notebook checkpoints folder 2024-07-22 12:03:25 -07:00
xander 006a3b9750 add ti config 2024-07-22 20:53:49 +02:00
aiXander b166f7c274 merge 2024-07-22 11:47:38 -07:00
aiXander f6cfd6d0bc fix typo 2024-07-22 11:46:33 -07:00
xander 6a4fa1ef8d add progressbar to node 2024-07-18 21:04:43 +02:00
xander 3cd17086f3 trainer comfyui v1 2024-07-18 05:50:04 +02:00
xander 75137a7539 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:07:07 +02:00
xander ee91f8c0c7 update naming conventions 2024-07-18 04:07:00 +02:00
xander 598c4204be add test config 2024-07-18 04:05:10 +02:00
xander 8da6c91af8 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:00:30 +02:00
xander 01523350e1 cleanup save dir 2024-07-18 04:00:21 +02:00
xander c891203365 update gitignore 2024-07-18 04:00:08 +02:00
xander 85599dc725 minor changes 2024-07-18 03:43:43 +02:00
xander af16d1edcc update configs 2024-07-18 03:26:51 +02:00
xander b3db4ff3bb Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-07-18 03:16:43 +02:00
xander 2a0761ae75 use gpt-4o 2024-07-18 03:16:41 +02:00
xander d8ee9ee522 use gpt-4o 2024-07-18 03:14:31 +02:00
xander 4f084d251c updates for comfyui node 2024-07-18 03:12:11 +02:00
aiXander fb521e9dbd add print flush 2024-07-16 04:39:29 -07:00
aiXander f937106f6c push changes 2024-07-16 04:04:34 -07:00
Gene Kogan 5f8fe5c4c8 Update requirements.txt 2024-07-16 03:44:59 -07:00
Gene Kogan 2f5eaeba7a Update cog.yaml 2024-07-16 03:44:35 -07:00
aiXander 7d8c845765 update yaml and print deps in config 2024-07-12 05:16:06 -07:00
aiXander 340336b53d print config pre training start 2024-07-11 06:11:33 -07:00
xander 5abed1487f setup lora_scale automation for validation grid 2024-07-11 13:36:11 +02:00
xander b5cf857b1f push small changes 2024-07-09 18:33:52 +02:00
xander e1813c75d1 updates to avoid OOM when plotting hist 2024-07-06 22:18:47 +02:00
xander 4fde1b7dd6 small tweaks 2024-07-06 18:45:38 +02:00
xander 892cf61c52 add disable_ti flag 2024-07-06 18:43:03 +02:00
xander fdd531d3c5 fix full finetuning and add 8bitadam 2024-07-06 17:42:49 +02:00
xander ef6635fafb update train cmd 2024-07-05 17:11:06 +02:00
xander faf864ce15 tiny tweak to learning rates for SDXL 2024-07-05 17:07:33 +02:00
xander b02221b23c cleanup before sd3 integration 2024-06-13 14:24:32 +02:00
xander d8b3175b57 cleanup before sd3 integration 2024-06-13 13:57:20 +02:00
50 changed files with 2799 additions and 1277 deletions
-7
View File
@@ -23,10 +23,3 @@ rendered_images
# trained models: # trained models:
lora_models/* lora_models/*
# Ignore the entire models folder by default:
models/*
### Include pipeline models: ###
!models/juggernaut_reborn.safetensors
!models/juggernaut_v6.safetensors
+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 }}
+3 -2
View File
@@ -1,22 +1,23 @@
cache cache
__pycache__ __pycache__
.ipynb_checkpoints/
models models
lora_models* lora_models*
eden_lora_training_runs/
datasets datasets
*.tar *.tar
.env .env
.cog .cog
.huggingface .huggingface
train.py
rendered_images* rendered_images*
gridsearch* gridsearch*
aesthetic_score_best_model.pth aesthetic_score_best_model.pth
# experiment folders: # experiment folders:
scripts/plots
conditioning_spaces/ conditioning_spaces/
training_args_x_*.json training_args_x_*.json
xander_configs/ 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.
+30 -50
View File
@@ -1,10 +1,11 @@
# Trainer # 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'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**! 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"> <p align="center">
<strong>Training images:</strong><br> <strong>Training images:</strong><br>
@@ -16,11 +17,26 @@ The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
</p> </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. - **style**: used for learning the aesthetic style of a collection of images.
- **face**: used for learning a specific face (can be human, character, ...). - **face**: used for learning a specific face (can be human, character, ...).
- **object**: will learn a specific object or thing featured in the training images. - **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 ## Setup
Install all dependencies using Install all dependencies using
@@ -29,7 +45,7 @@ Install all dependencies using
then you can simply run: then you can simply run:
`python main.py -c training_args.json` `python main.py train_configs/training_args.json`
to start a training job. to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job. Adjust the arguments inside `training_args.json` to setup a custom training job.
@@ -44,61 +60,25 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
sudo chmod +x /usr/local/bin/cog sudo chmod +x /usr/local/bin/cog
``` ```
2. Build the image with `sudo cog build` 2. Build the image with `cog build`
3. Run a training run with `sudo sh cog_test_train.sh` 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 ## TODO's
Bugs: 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! - 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!
- figure out why training is 3x slower through comfyui node versus just running main.py as a python job..?
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)
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts) - 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: Bigger improvements:
- add stronger token regularization (eg CelebBasis spanning basis): - integrate Flux / SD3
- Add multi-token training - 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 perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- implement prompt-aligned: https://prompt-aligned.github.io/ - 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

+4 -7
View File
@@ -3,17 +3,14 @@
build: build:
gpu: true gpu: true
cuda: "11.8" cuda: "12.1"
python_version: "3.9" python_version: "3.11"
system_packages: system_packages:
- "ffmpeg" - "ffmpeg"
- "libgl1-mesa-glx"
- "libegl1-mesa-dev"
- "libsm6"
- "libxext6"
python_requirements: requirements.txt python_requirements: requirements.txt
run: run:
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task - wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
predict: "predict.py:Predictor" predict: "predict.py:Predictor"
image: "r8.im/edenartlab/sdxl-lora-trainer" image: "r8.im/edenartlab/sdxl-lora-trainer"
+7 -6
View File
@@ -1,12 +1,13 @@
# Set GPU ID to run these jobs on: # Set GPU ID to run these jobs on:
GPU_ID="device=2" GPU_ID="device=3"
cog predict --gpus $GPU_ID \ cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \ -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 lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
-i concept_mode="face" \ -i concept_mode="style" \
-i sd_model_version="sdxl" \ -i sd_model_version="sdxl" \
-i max_train_steps="360" \ -i max_train_steps="300" \
-i caption_model="blip" \ -i sample_imgs_lora_scale="0.7" \
-i debug="False" \ -i n_sample_imgs="6" \
-i debug="True" \
-i seed="0" -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)}')
+131 -110
View File
@@ -1,4 +1,3 @@
import fnmatch
import math import math
import os import os
import time import time
@@ -12,20 +11,18 @@ import torch
import torch.utils.checkpoint import torch.utils.checkpoint
from tqdm import tqdm from tqdm import tqdm
import prodigyopt
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.utils.utils import *
from trainer.checkpoint import save_checkpoint from trainer.checkpoint import save_checkpoint
from trainer.embedding_handler import TokenEmbeddingsHandler from trainer.embedding_handler import TokenEmbeddingsHandler
from trainer.dataset import PreprocessedDataset from trainer.dataset import PreprocessedDataset
from trainer.config import TrainingConfig from trainer.config import TrainingConfig
from trainer.models import print_trainable_parameters, load_models 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.inference import render_images, get_conditioning_signals
from trainer.preprocess import preprocess from trainer.preprocess import preprocess
from trainer.utils.io import make_validation_img_grid 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 ( from trainer.optimizer import (
OptimizerCollection, OptimizerCollection,
get_optimizer_and_peft_models_text_encoder_lora, get_optimizer_and_peft_models_text_encoder_lora,
@@ -34,10 +31,43 @@ from trainer.optimizer import (
get_unet_optimizer get_unet_optimizer
) )
def train( def train(config: TrainingConfig):
config: TrainingConfig,
):
seed_everything(config.seed) seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
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, input_dir = preprocess(
config, config,
@@ -58,19 +88,6 @@ def train(
if config.allow_tf32: if config.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cuda.matmul.allow_tf32 = True
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(config.pretrained_model, config.device, weight_dtype, keep_vae_float32=0)
# Initialize new tokens for training. # Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler( embedding_handler = TokenEmbeddingsHandler(
text_encoders = [text_encoder_one, text_encoder_two], text_encoders = [text_encoder_one, text_encoder_two],
@@ -113,22 +130,26 @@ def train(
embedding_handler.make_embeddings_trainable() embedding_handler.make_embeddings_trainable()
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer( if not config.disable_ti:
text_encoders=text_encoders, optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
textual_inversion_lr=config.ti_lr, text_encoders=text_encoders,
textual_inversion_weight_decay=config.ti_weight_decay, textual_inversion_lr=config.ti_lr,
optimizer_name=config.ti_optimizer ## hardcoded textual_inversion_weight_decay=config.ti_weight_decay,
) optimizer_name=config.ti_optimizer ## hardcoded
)
else:
optimizer_ti = None
textual_inversion_params = None
if not config.is_lora: # This code pathway has not been tested in a long while if not config.is_lora: # This code pathway has not been tested in a long while
print(f"Doing full fine-tuning on the U-Net") print(f"Doing full fine-tuning on the U-Net")
unet.requires_grad_(True) unet.requires_grad_(True)
unet_lora_parameters = None unet_lora_parameters = None
optimizer_text_encoder_lora = None
unet_trainable_params = unet.parameters() unet_trainable_params = unet.parameters()
else: else:
# Do lora-training instead. # Do lora-training instead.
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora # https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
# target_blocks=["block"] for original IP-Adapter # target_blocks=["block"] for original IP-Adapter
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only # target_blocks=["up_blocks.0.attentions.1"] for style blocks only
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks # target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
@@ -142,14 +163,17 @@ def train(
pipe=pipe pipe=pipe
) )
optimizer_unet = get_unet_optimizer( if config.unet_lr > 0.0:
prodigy_d_coef=config.prodigy_d_coef, optimizer_unet = get_unet_optimizer(
prodigy_growth_factor=config.unet_prodigy_growth_factor, prodigy_d_coef=config.prodigy_d_coef,
lora_weight_decay=config.lora_weight_decay, prodigy_growth_factor=config.unet_prodigy_growth_factor,
use_dora=config.use_dora, lora_weight_decay=config.lora_weight_decay,
unet_trainable_params=unet_trainable_params, use_dora=config.use_dora,
optimizer_name=config.unet_optimizer_type unet_trainable_params=unet_trainable_params,
) optimizer_name=config.unet_optimizer_type
)
else:
optimizer_unet = None
print_trainable_parameters(unet, model_name = 'unet') print_trainable_parameters(unet, model_name = 'unet')
for i, text_encoder in enumerate(text_encoders): for i, text_encoder in enumerate(text_encoders):
@@ -161,31 +185,26 @@ def train(
pipe, pipe,
vae.float(), vae.float(),
size = config.train_img_size, size = config.train_img_size,
do_cache=config.do_cache,
substitute_caption_map=config.token_dict, substitute_caption_map=config.token_dict,
aspect_ratio_bucketing=config.aspect_ratio_bucketing, aspect_ratio_bucketing=config.aspect_ratio_bucketing,
train_batch_size=config.train_batch_size 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') vae = vae.to('cpu')
gc.collect() gc.collect()
torch.cuda.empty_cache() torch.cuda.empty_cache()
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
train_dataloader = torch.utils.data.DataLoader( train_dataloader = torch.utils.data.DataLoader(
train_dataset, train_dataset,
batch_size=config.train_batch_size, batch_size=config.train_batch_size,
shuffle=True, 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) config.num_train_epochs = int(math.ceil(config.max_train_steps / len(train_dataloader)))
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
if config.max_train_steps is None:
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
print(f"--- Num samples = {len(train_dataset)}") print(f"--- Num samples = {len(train_dataset)}")
@@ -194,7 +213,7 @@ def train(
print(f"--- Instantaneous batch size per device = {config.train_batch_size}") print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}") print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}") print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
print(f"--- Total optimization steps = {config.max_train_steps}\n") print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
global_step = 0 global_step = 0
last_save_step = 0 last_save_step = 0
@@ -208,19 +227,17 @@ def train(
# Data tracking inits: # Data tracking inits:
start_time, images_done = time.time(), 0 start_time, images_done = time.time(), 0
prompt_embeds_norms = {'main':[], 'reg':[]} 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': []}, {} grad_norms, token_stds = {'unet': []}, {}
for i in range(len(text_encoders)): for i in range(len(text_encoders)):
grad_norms[f'text_encoder_{i}'] = [] grad_norms[f'text_encoder_{i}'] = []
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)} token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
# default value of cold (pre-warmup) optimizer lr: # default values for cold (starting) optimizer lr:
if config.sd_model_version == "sdxl": base_unet_lr = 2.0e-4 if (config.is_lora and config.disable_ti) else 5.0e-5
# let textual_inversion do the work first!
base_lr = 0.5e-5 if not config.is_lora:
elif config.sd_model_version == "sd15": base_unet_lr = 1.0e-5
# let lora training kick in soonish
base_lr = 1.0e-4
####################################################################################################### #######################################################################################################
@@ -235,7 +252,8 @@ def train(
) )
optimizers = optimizer_collection.optimizers 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): for epoch in range(config.num_train_epochs):
if config.aspect_ratio_bucketing: if config.aspect_ratio_bucketing:
@@ -248,15 +266,12 @@ def train(
completion_f = finegrained_epoch / config.num_train_epochs completion_f = finegrained_epoch / config.num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training # 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 config.ti_optimizer != "prodigy" and optimizers['textual_inversion'] is not None:
if optimizers['textual_inversion'] is not None: # Apply the exponential learning rate
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0 optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 1.7
# warmup the ti-lr: # Apply freezing condition
if config.ti_lr_warmup_steps > 0: if completion_f > config.freeze_ti_after_completion_f:
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0) optimizers['textual_inversion'].param_groups[0]['lr'] = 0.0
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
if config.freeze_ti_after_completion_f <= completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
if optimizers['text_encoders'] is not None: if optimizers['text_encoders'] is not None:
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0 optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
@@ -268,16 +283,26 @@ def train(
if optimizers['unet'] is not None: if optimizers['unet'] is not None:
# Calculate the exponential factor # 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 # 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: if not config.aspect_ratio_bucketing:
captions, vae_latent, mask = batch captions, vae_latent, mask = batch
else: else:
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch() captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
mask = mask.to(config.device)
captions = list(captions) 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( prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
config, pipe, captions config, pipe, captions
) )
@@ -314,13 +339,18 @@ def train(
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps) loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
losses['img_loss'].append(loss.item()) 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: 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) 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: # Dont apply this loss, just plot it for now:
loss += 0.0 * concept_description_loss loss += 0.0 * concept_description_loss
losses['concept_description_loss'].append(concept_description_loss.item()) losses['concept_description_loss'].append(concept_description_loss.item())
if config.l1_penalty > 0.0: if config.l1_penalty > 0.0 and unet_lora_parameters:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters: # Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters) l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += config.l1_penalty * l1_norm loss += config.l1_penalty * l1_norm
@@ -349,12 +379,6 @@ def train(
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm) grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
optimizer_collection.step() optimizer_collection.step()
# after every optimizer step, we do some manual intervention of the embeddings to regularize them:
if optimizer_collection.get_lr('textual_inversion') > 0.0:
#embedding_handler.fix_embedding_std(config.off_ratio_power)
pass
optimizer_collection.zero_grad() optimizer_collection.zero_grad()
############################################################################################################# #############################################################################################################
@@ -368,8 +392,13 @@ def train(
for std_i, std in enumerate(embedding_stds): for std_i, std in enumerate(embedding_stds):
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item()) 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: # Print some statistics:
if config.debug and (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > -1: 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}" output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True) os.makedirs(output_save_dir, exist_ok=True)
@@ -390,26 +419,16 @@ def train(
) )
last_save_step = global_step last_save_step = global_step
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings() if config.debug:
for idx, text_encoder in enumerate(text_encoders): embedding_handler.print_token_info()
if text_encoder is None: if config.is_lora: # plotting this hist for full unet parameters can run OOM
continue 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)
n = len(token_embeddings[f'txt_encoder_{idx}']) plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
for i in range(n): target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
token = trainable_tokens[f'txt_encoder_{idx}'][i] plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
# Strip any backslashes from the token name: plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
token = token.replace("/", "_") plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
embedding = token_embeddings[f'txt_encoder_{idx}'][i] #plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
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()
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
validation_prompts = render_images( validation_prompts = render_images(
pipe = pipe, pipe = pipe,
@@ -420,6 +439,8 @@ def train(
is_lora = config.is_lora, is_lora = config.is_lora,
pretrained_model = config.pretrained_model, pretrained_model = config.pretrained_model,
lora_scale = config.sample_imgs_lora_scale, lora_scale = config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs, n_imgs = config.n_sample_imgs,
device = config.device, device = config.device,
checkpoint_folder = None checkpoint_folder = None
@@ -433,14 +454,13 @@ def train(
images_done += config.train_batch_size images_done += config.train_batch_size
global_step += 1 global_step += 1
if global_step % (config.max_train_steps//20) == 0: if global_step % (config.max_train_steps//100) == 0:
progress = (global_step / config.max_train_steps) + 0.05 progress = (global_step / config.max_train_steps) + 0.05
print_system_info() #print_system_info()
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
yield np.min((progress, 1.0)) yield np.min((progress, 1.0))
if global_step > config.max_train_steps: if global_step > config.max_train_steps:
print("Reached max steps, stopping training!") print("Reached max steps, stopping training!", flush = True)
break break
# final_save # final_save
@@ -471,8 +491,7 @@ def train(
pretrained_model_version=config.pretrained_model["version"] pretrained_model_version=config.pretrained_model["version"]
) )
print("Running final inference round...") if config.debug and 0:
if config.debug:
# Reload the entire pipe from disk + LoRa: # Reload the entire pipe from disk + LoRa:
pipe_to_use = None pipe_to_use = None
checkpoint_folder = output_save_dir checkpoint_folder = output_save_dir
@@ -502,6 +521,8 @@ def train(
is_lora=config.is_lora, is_lora=config.is_lora,
pretrained_model=config.pretrained_model, pretrained_model=config.pretrained_model,
lora_scale=config.sample_imgs_lora_scale, lora_scale=config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs, n_imgs = config.n_sample_imgs,
n_steps = 30, n_steps = 30,
device = config.device, device = config.device,
@@ -511,13 +532,6 @@ def train(
img_grid_path = make_validation_img_grid(output_save_dir) img_grid_path = make_validation_img_grid(output_save_dir)
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg")) shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
# Remove unneeded checkpoints if they exist in the output directory:
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_save_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
else: else:
print(f"Skipping final save, {output_save_dir} already exists") print(f"Skipping final save, {output_save_dir} already exists")
@@ -531,6 +545,8 @@ def train(
config.job_time = time.time() - config.start_time config.job_time = time.time() - config.start_time
config.training_attributes["validation_prompts"] = validation_prompts config.training_attributes["validation_prompts"] = validation_prompts
config.save_as_json(os.path.join(output_save_dir, "training_args.json")) config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
print("Training job complete, saving outputs...", flush = True)
print("------------------------------------------")
return config, output_save_dir return config, output_save_dir
@@ -541,6 +557,11 @@ if __name__ == "__main__":
args = parser.parse_args() args = parser.parse_args()
config = TrainingConfig.from_json(file_path=args.config_filename) config = TrainingConfig.from_json(file_path=args.config_filename)
print("Starting new LoRa training run with config:")
print(config)
print("------------------------------------------")
for progress in train(config=config): for progress in train(config=config):
print(f"Progress: {(100*progress):.2f}%", end="\r") print(f"Progress: {(100*progress):.2f}%", end="\r")
+93 -60
View File
@@ -1,22 +1,17 @@
import os import os
import shutil
import tarfile import tarfile
import json import json
import time import time
import random
import torch import torch
import numpy as np import numpy as np
import pandas as pd from PIL import Image
from dotenv import load_dotenv
from main import train from main import train
from trainer.config import TrainingConfig, model_paths
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.io import clean_filename
from trainer.utils.utils import seed_everything
import folder_paths
import comfy.utils
class Eden_LoRa_trainer: class Eden_LoRa_trainer:
@classmethod @classmethod
@@ -24,92 +19,130 @@ class Eden_LoRa_trainer:
return { return {
"required": { "required": {
"training_images_folder_path": ("STRING", {"default": "."}), "training_images_folder_path": ("STRING", {"default": "."}),
"lora_name": ("STRING", {"default": ""}), "mode": (["style", "face", "object"], {"default": "style"}),
"sd_model_version": (["sdxl", "sd15"], ), "lora_name": ("STRING", {"default": "Eden_Token_LoRa"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}), "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"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}), "train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}), "max_train_steps": ("INT", {"default": 300, "min": 10, "max": 10000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}), "ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0, "max": 0.005, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "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}), "lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
"use_dora": ("BOOLEAN", {"default": False}), "disable_ti": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}), "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}),
} }
} }
CATEGORY = "Eden 🌱" CATEGORY = "Eden 🌱"
RETURN_TYPES = ("STRING",) RETURN_TYPES = ("IMAGE", "STRING", "STRING", "STRING")
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
FUNCTION = "train_lora" FUNCTION = "train_lora"
def train_lora(self, training_images_folder_path, def train_lora(self,
name = lora_name, training_images_folder_path,
concept_mode = "style", ckpt_name,
sd_model_version = "sdxl", lora_name,
seed = 0, mode,
resolution = 521, training_resolution,
train_batch_size = 4, train_batch_size,
max_train_steps = 400, max_train_steps ,
ti_lr = 0.001, ti_lr,
unet_lr = 0.001, unet_lr,
lora_rank = 16, lora_rank,
use_dora = False, disable_ti,
n_tokens = 2 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...") 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"))
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
config = TrainingConfig( config = TrainingConfig(
name="test", name=lora_name,
output_dir="output",
lora_training_urls=training_images_folder_path, lora_training_urls=training_images_folder_path,
concept_mode=concept_mode, concept_mode=mode,
sd_model_version=sd_model_version, ckpt_path=ckpt_path,
seed=seed, seed=seed,
resolution=resolution, resolution=training_resolution,
train_batch_size=train_batch_size, train_batch_size=train_batch_size,
max_train_steps=max_train_steps, max_train_steps=max_train_steps,
checkpointing_steps=10000, 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, ti_lr=ti_lr,
unet_lr=unet_lr, unet_lr=unet_lr,
lora_rank=lora_rank, lora_rank=lora_rank,
use_dora=use_dora, use_dora=False,
caption_model="blip", caption_model="blip",
disable_ti=disable_ti,
n_tokens=n_tokens, n_tokens=n_tokens,
verbose=True, verbose=True,
debug=True, debug=plot_training_graphs_on_disk,
) )
pbar = comfy.utils.ProgressBar(100)
with torch.inference_mode(False): with torch.inference_mode(False):
train_generator = train(config=config) train_generator = train(config=config)
while True: while True:
try: try:
progress_f = next(train_generator) progress_f = next(train_generator)
pbar.update_absolute(progress_f * 100)
except StopIteration as e: except StopIteration as e:
config, output_save_dir = e.value # Capture the return value config, output_save_dir = e.value # Capture the return value
break break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
out_path = f"{clean_filename(lora_name)}_eden_concept_lora_{int(time.time())}.tar"
directory = cogPath(output_save_dir)
with tarfile.open(out_path, "w") as tar:
print("Adding files to tar...")
for file_path in directory.rglob("*"):
print(file_path)
arcname = file_path.relative_to(directory)
tar.add(file_path, arcname=arcname)
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
attributes = {} attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"] attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time attributes['job_time_seconds'] = config.job_time
print(f"LORA training finished in {config.job_time:.1f} seconds") print(f"LORA training node finished in {config.job_time:.1f} seconds")
print(f"Returning {out_path}") print("---------- Made with love by Eden.art 🌱 ----------")
return (out_path,) # safetensors paths:
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
# find the index of the path containing "_embeddings.safetensors":
for i, path in enumerate(paths):
if "_embeddings.safetensors" in path:
embedding_path = path
else:
lora_path = path
# Load the grid 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_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 typing import Iterator, Optional
from trainer.preprocess import preprocess from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig from trainer.config import TrainingConfig
from trainer.utils.io import clean_filename from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything from trainer.utils.utils import seed_everything
@@ -66,19 +65,19 @@ class Predictor(BasePredictor):
), ),
max_train_steps: int = Input( 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", 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( 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 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( unet_lr: float = Input(
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting", 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( ti_lr: float = Input(
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.", 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.", description="Rank of LoRA embeddings for the unet.",
default=16 default=16
), ),
use_dora: bool = Input(
description="Use Dora instead of LoRa",
default=False,
),
n_tokens: int = Input( n_tokens: int = Input(
description="How many new tokens to train (highly recommended to leave this at 2)", 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( seed: int = Input(
description="Random seed for reproducible training. Leave empty to use a random seed", 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, sd_model_version=sd_model_version,
seed=seed, seed=seed,
resolution=resolution, 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, train_batch_size=train_batch_size,
max_train_steps=max_train_steps, max_train_steps=max_train_steps,
checkpointing_steps=10000, checkpointing_steps=checkpointing_steps,
n_sample_imgs=n_sample_imgs,
ti_lr=ti_lr, ti_lr=ti_lr,
unet_lr=unet_lr, unet_lr=unet_lr,
lora_rank=lora_rank, lora_rank=lora_rank,
use_dora=use_dora,
caption_model="blip", caption_model="blip",
n_tokens=n_tokens, n_tokens=n_tokens,
verbose=True, verbose=True,
@@ -162,9 +175,13 @@ class Predictor(BasePredictor):
# Add instructions README: # Add instructions README:
tar.add("instructions_README.md", arcname="README.md") tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json") comfy_workflows_path = "ComfyUI_workflows"
if sd_model_version == "sd15": if os.path.exists(comfy_workflows_path) and os.path.isdir(comfy_workflows_path):
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json") 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 = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"] 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 = ""
+23 -14
View File
@@ -1,17 +1,26 @@
torch>=2.1.0 torch>=2.1.0
torchaudio>=2.1.0
torchvision>=0.16.0 torchvision>=0.16.0
transformers>=4.38.1 transformers>=4.38.0
diffusers>=0.27.2 diffusers>=0.29.2
ujson>=5.9.0 tokenizers>=0.15.2
scipy>=1.12.0 huggingface-hub==0.23.2
peft>=0.10.0 ujson==5.10.0
invisible-watermark>=0.2.0 scipy==1.14.0
peft==0.10.0
invisible-watermark==0.2.0
pandas==2.2.1 pandas==2.2.1
numpy>=1.26.4 numpy==1.26.4
opencv-python>=4.1.0.25 opencv-python==4.10.0.84
mediapipe>=0.10.11 mediapipe==0.10.14
openai>=1.14.0 openai==1.35.13
python-dotenv python-dotenv==1.0.1
prodigyopt prodigyopt==1.0
omegaconf omegaconf==2.3.0
ujson 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
+40 -38
View File
@@ -1,10 +1,8 @@
""" """
Faces: 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.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/gene.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
Objects: Objects:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
@@ -14,7 +12,12 @@ https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets
Styles: 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/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: # Setup the base experiment config:
exp_name = "grimes" exp_name = "ygor_sd15"
caption_prefix = "" caption_prefix = ""
mask_target_prompts = "" mask_target_prompts = ""
n_exp = 200 # how many random experiment settings to generate n_exp = 100 # how many random experiment settings to generate
min_hamming_distance = 3 # min_n_params that have to be different from any previous experiment to be scheduled 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" output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values # Define training hyperparameters and their possible values
@@ -48,43 +51,40 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
hyperparameters = { hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"], "output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sd15", "sdxl"], "sd_model_version": ["sd15"],
"lora_training_urls": [ "lora_training_urls": [
"/home/rednax/Documents/datasets/grimes" "/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean"
], ],
"concept_mode": ['face'], "concept_mode": ['style'],
"sample_imgs_lora_scale": [0.9],
"caption_dropout": [0.2],
"seed": [0], "seed": [0],
"resolution": [512], "resolution": [512,640,768],
"train_batch_size": [4], "train_batch_size": [8],
"n_sample_imgs": [6], "n_sample_imgs": [8],
"max_train_steps": [400,800], "max_train_steps": [2000],
"checkpointing_steps": [100], "checkpointing_steps": [500],
"gradient_accumulation_steps": [1], "gradient_accumulation_steps": [1],
"n_tokens": [2], "n_tokens": [3],
"ti_lr": [0.001,0.0005], "disable_ti": ['true'],
"ti_weight_decay": [0.001,0.0], "ti_lr": [0.0001],
"l1_penalty": [0.0], "token_warmup_steps": [0],
"token_warmup_steps": [0,60],
"tok_cov_reg_w": [2000],
"cond_reg_w": [0.01e-5],
"tok_cond_reg_w": [0.01e-5],
"unet_prodigy_growth_factor": [1.05], "unet_lr": [0.001, 0.0003],
"unet_lr": [0.001], "lora_rank": [8,24,64],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16,32],
"use_dora": ['false', 'true'], "use_dora": ['false', 'true'],
"unet_optimizer_type": ['adamw'],
"is_lora": ['true'],
"text_encoder_lora_optimizer": [None], "text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4], "text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0], "snr_gamma": [5.0],
"caption_model": ["blip", "gpt4-v"], "caption_model": ["florence", "blip", "no_caption"],
"augment_imgs_up_to_n": [20,40], "augment_imgs_up_to_n": [40],
"verbose": ['true'], "verbose": ['true'],
"debug": ['true'] "debug": ['true']
} }
@@ -100,7 +100,7 @@ shutil.rmtree(config_output_dir, ignore_errors=True)
os.makedirs(config_output_dir, exist_ok=True) os.makedirs(config_output_dir, exist_ok=True)
# Open the shell script file # 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 for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
resamples, combination = 0, None resamples, combination = 0, None
@@ -122,9 +122,6 @@ for exp_index in tqdm(range(n_exp)): # number of combinations you want to gener
dirname = os.path.dirname(config_filename) dirname = os.path.dirname(config_filename)
os.makedirs(dirname, exist_ok=True) 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: with open(config_filename, "w") as f:
json.dump(experiment_settings, f, indent=4) json.dump(experiment_settings, f, indent=4)
break break
@@ -146,7 +143,12 @@ def generate_sh_script(folder_path, output_sh_path):
# Write a command for each JSON file # Write a command for each JSON file
for json_file in json_files: for json_file in json_files:
command = f"python main.py {os.path.join(folder_path, json_file)}\n" file_path = os.path.join("scripts/", folder_path, json_file)
command = f"python main.py {file_path}\n"
if nohup:
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
sh_file.write(command) sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path) generate_sh_script(config_output_dir, output_sh_path)
+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)
@@ -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
}
+20
View File
@@ -0,0 +1,20 @@
{
"unet_lr": 0.00,
"lora_rank": 4,
"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
}
@@ -0,0 +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.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 768,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 600,
"checkpointing_steps": 200,
"disable_ti": false,
"caption_model": "florence",
"ti_lr": 0.001,
"unet_lr": 0.0005,
"lora_rank": 16,
"debug": true
}
@@ -0,0 +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.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": 300,
"checkpointing_steps": 200,
"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
}
@@ -0,0 +1,21 @@
{
"name": "twisting_realities_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.8,
"seed": 0,
"resolution": 768,
"train_batch_size": 4,
"n_sample_imgs": 8,
"max_train_steps": 1200,
"checkpointing_steps": 200,
"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
}
@@ -0,0 +1,21 @@
{
"name": "twisting_realities",
"sd_model_version": "sdxl",
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.75,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"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
}
+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
}
+36 -5
View File
@@ -54,10 +54,31 @@ def set_adapter_scales(pipe, lora_scale = 1.0):
return pipe 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
def remove_delimiter_characters(name: str): # Replace all occurrences of the pattern with a single underscore
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath: cleaned_name = re.sub(pattern, '_', name)
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
# 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]
# 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 # Convert to WebUI format
def convert_pytorch_lora_safetensors_to_webui( def convert_pytorch_lora_safetensors_to_webui(
@@ -135,7 +156,7 @@ def save_checkpoint(
embedding_handler.save_embeddings( embedding_handler.save_embeddings(
os.path.join( os.path.join(
output_dir, output_dir,
f"{name}_embeddings.safetensors" f"{name}_{pretrained_model_version}_embeddings.safetensors"
) )
) )
@@ -184,11 +205,21 @@ def save_checkpoint(
convert_pytorch_lora_safetensors_to_webui( convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"), pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
output_filename=os.path.join(output_dir, f"{name}.safetensors") output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_lora.safetensors")
) )
else: else:
# Save the entire, finetuned unet weights:
unet.save_pretrained(save_directory = output_dir) unet.save_pretrained(save_directory = output_dir)
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
return
def load_checkpoint( def load_checkpoint(
pretrained_model_version: str, pretrained_model_version: str,
pretrained_model_path: str, pretrained_model_path: str,
+66 -33
View File
@@ -3,15 +3,47 @@ from datetime import datetime
from pydantic import BaseModel from pydantic import BaseModel
import json, time, os import json, time, os
from typing import Literal from typing import Literal
from trainer.models import pretrained_models
from trainer.utils.utils import pick_best_gpu_id 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",
}
def get_path(self, key):
return self.paths.get(key, None)
def set_path(self, key, path):
if key in self.paths:
self.paths[key] = path
model_paths = ModelPaths()
# Default 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"
SD15_URL = "https://huggingface.co/KamCastle/jugg/resolve/main/juggernaut_reborn.safetensors"
pretrained_models = {
"sdxl": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SDXL_URL)), "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SD15_URL)), "url": SD15_URL, "version": "sd15"}
}
class TrainingConfig(BaseModel): class TrainingConfig(BaseModel):
lora_training_urls: str lora_training_urls: str
concept_mode: Literal["face", "style", "object"] 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_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
sd_model_version: Literal["sdxl", "sd15"] 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 pretrained_model: dict = None
seed: Union[int, None] = None seed: Union[int, None] = None
resolution: int = 512 resolution: int = 512
@@ -19,36 +51,36 @@ class TrainingConfig(BaseModel):
train_img_size: List[int] = None train_img_size: List[int] = None
train_aspect_ratio: float = None train_aspect_ratio: float = None
train_batch_size: int = 4 train_batch_size: int = 4
num_train_epochs: int = 10000 max_train_steps: int = 300
max_train_steps: int = 360 num_train_epochs: int = None
checkpointing_steps: int = 10000 checkpointing_steps: int = 10000
gradient_accumulation_steps: int = 1 gradient_accumulation_steps: int = 1
is_lora: bool = True is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy"] = "adamw" 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_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 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) 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: float = 0.001
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0 ti_weight_decay: float = 0.0
ti_optimizer: Literal["adamw", "prodigy"] = "adamw" 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 cond_reg_w: float = 0.0e-5
tok_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 tok_cov_reg_w: float = 0. # regularizes the token covariance matrix wrt pretrained, normal tokens
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std l1_penalty: float = 0.03 # Makes the unet lora matrix more sparse
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
snr_gamma: float = 5.0 snr_gamma: float = 5.0
lora_alpha_multiplier: float = 1.0 lora_alpha_multiplier: float = 1.0
lora_rank: int = 12 lora_rank: int = 16
use_dora: bool = False use_dora: bool = False
left_right_flip_augmentation: bool = True left_right_flip_augmentation: bool = True
@@ -59,22 +91,17 @@ class TrainingConfig(BaseModel):
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
n_sample_imgs: int = 4 n_sample_imgs: int = 4
name: str = None name: str = None
output_dir: str = "lora_models/unnamed" output_dir: str = "eden_lora_training_runs"
debug: bool = False debug: bool = False
allow_tf32: bool = True allow_tf32: bool = True
remove_ti_token_from_prompts: bool = False disable_ti: bool = False
skip_gpt_cleanup: bool = False
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16" weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 2 n_tokens: int = 3
inserting_list_tokens: List[str] = ["<s0>","<s1>"] inserting_list_tokens: List[str] = ["<s0>","<s1>","<s2>"]
token_dict: dict = {"TOK": "<s0><s1>"} token_dict: dict = {"TOK": "<s0><s1><s2>"}
device: str = "cuda:0" device: str = "cuda:0"
crops_coords_top_left_h: int = 0 sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
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 = 0.65 # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0 dataloader_num_workers: int = 0
training_attributes: dict = {} training_attributes: dict = {}
aspect_ratio_bucketing: bool = False aspect_ratio_bucketing: bool = False
@@ -93,16 +120,19 @@ class TrainingConfig(BaseModel):
def __init__(self, **data): def __init__(self, **data):
super().__init__(**data) super().__init__(**data)
self.pretrained_model = pretrained_models[self.sd_model_version]
# add some metrics to the foldername: if not self.ckpt_path:
lora_str = "dora" if self.use_dora else "lora" self.pretrained_model = pretrained_models[self.sd_model_version]
timestamp_short = datetime.now().strftime("%d_%H-%M-%S") else:
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
if not self.name: 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"--{timestamp_short}-{self.sd_model_version}_{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) os.makedirs(self.output_dir, exist_ok=True)
if self.seed is None: if self.seed is None:
@@ -111,6 +141,9 @@ class TrainingConfig(BaseModel):
if self.unet_lr_warmup_steps is None: if self.unet_lr_warmup_steps is None:
self.unet_lr_warmup_steps = self.max_train_steps 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": if self.concept_mode == "face":
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to '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.left_right_flip_augmentation = False # always disable lr flips for face mode!
+35 -31
View File
@@ -1,6 +1,7 @@
import os import os
import torch import torch
import numpy as np import numpy as np
from tqdm import tqdm
import pandas as pd import pandas as pd
import PIL import PIL
from PIL import Image from PIL import Image
@@ -32,7 +33,6 @@ class PreprocessedDataset(Dataset):
data_dir: str, data_dir: str,
pipe, pipe,
vae_encoder, vae_encoder,
do_cache: bool = False,
size: List[int] = [512, 512], size: List[int] = [512, 512],
text_dropout: float = 0.0, text_dropout: float = 0.0,
aspect_ratio_bucketing: bool = False, aspect_ratio_bucketing: bool = False,
@@ -42,13 +42,14 @@ class PreprocessedDataset(Dataset):
super().__init__() super().__init__()
self.data_dir = data_dir self.data_dir = data_dir
self.csv_path = os.path.join(data_dir, "captions.csv") 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.data["caption"]
self.captions = self.captions.str.lower() self.captions = self.captions.str.lower()
for key, value in substitute_caption_map.items(): for key, value in substitute_caption_map.items():
self.captions = self.captions.str.replace(key.lower(), value) self.captions = self.captions.str.replace(key.lower(), value)
self.captions = self.captions.fillna("")
self.image_path = self.data["image_path"] self.image_path = self.data["image_path"]
if "mask_path" not in self.data.columns: if "mask_path" not in self.data.columns:
@@ -62,23 +63,31 @@ class PreprocessedDataset(Dataset):
self.text_dropout = text_dropout self.text_dropout = text_dropout
self.size = size self.size = size
if do_cache: # If the training data is small we can keep everything in memory, otherwise offload to disk
print("Caching latents, masks and captions...\n") 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.vae_latents = []
self.masks = [] self.masks = []
self.do_cache = True
for idx in range(len(self.data)): for idx in tqdm(range(len(self.data))):
if len(self.data) < 25: vae_latent, mask, _ = self._process(idx)
print(self.captions[idx])
vae_latent, mask = self._process(idx)
self.vae_latents.append(vae_latent) 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.") else: # Store the latents and masks on disk
del self.vae_encoder print("Encoding latents, masks and captions and storing on disk...\n")
else: self.vae_latents = None
self.do_cache = False 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: if aspect_ratio_bucketing:
print("Using aspect ratio bucketing.") print("Using aspect ratio bucketing.")
@@ -100,12 +109,9 @@ class PreprocessedDataset(Dataset):
def get_aspect_ratio_bucketed_batch(self): 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__()" 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() indices, resolution = self.bucket_manager.get_batch()
print(f"Got bucket batch: {indices}, resolution: {resolution}")
tok1, tok2, vae_latents, masks = [], [], [], [] tok1, tok2, vae_latents, masks = [], [], [], []
for idx in indices: for idx in indices:
if self.tokenizer_2 is None: if self.tokenizer_2 is None:
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution) t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
else: else:
@@ -141,28 +147,24 @@ class PreprocessedDataset(Dataset):
image = PIL.Image.open(image_path).convert("RGB") image = PIL.Image.open(image_path).convert("RGB")
if bucketing_resolution is None: if bucketing_resolution is None:
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to( 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: else:
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to( 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() dummy_vae_latent = vae_latent.sample()
if self.mask_path is None: if self.mask_path is None:
mask = torch.ones_like( mask = torch.ones_like(dummy_vae_latent, dtype=self.vae_encoder.dtype)
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
else: else:
mask_path = self.mask_path[idx] mask_path = self.mask_path[idx]
mask_path = os.path.join(self.data_dir, mask_path) mask_path = os.path.join(self.data_dir, mask_path)
mask = PIL.Image.open(mask_path) mask = PIL.Image.open(mask_path)
mask = prepare_mask(mask, self.size[0], self.size[1]).to( mask = prepare_mask(mask, self.size[0], self.size[1]).to(dtype=self.vae_encoder.dtype)
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
mask_dtype = mask.dtype mask_dtype = mask.dtype
mask = mask.float() mask = mask.float()
@@ -174,7 +176,7 @@ class PreprocessedDataset(Dataset):
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4 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__( def __getitem__(
self, idx: int, bucketing_resolution:tuple = None self, idx: int, bucketing_resolution:tuple = None
@@ -182,10 +184,12 @@ class PreprocessedDataset(Dataset):
if self.do_cache: if self.do_cache:
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
return self.captions[idx], vae_latent.squeeze(), self.masks[idx] return self.captions[idx], vae_latent.squeeze().detach(), self.masks[idx].detach()
else: # This code pathway has not been tested in a long time and might be broken else: # Load from disk:
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution) vae_latent = torch.load(os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
vae_latent = vae_latent.sample() * self.vae_scaling_factor 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 = (config.resolution, config.resolution)
original_size = (1024, 1024) original_size = (1024, 1024)
target_size = (config.resolution, config.resolution) target_size = (config.resolution, config.resolution)
crops_coords_top_left = (0,0)
crops_coords_top_left = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None: if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
@@ -397,7 +393,6 @@ class TokenEmbeddingsHandler:
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0. embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
optimizer_ti.step() optimizer_ti.step()
self.fix_embedding_std(config.off_ratio_power)
optimizer_ti.zero_grad() optimizer_ti.zero_grad()
if config.debug: if config.debug:
@@ -430,41 +425,6 @@ class TokenEmbeddingsHandler:
def device(self): def device(self):
return self.text_encoders[0].device 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): def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
# Assuming new tokens are of the format <s_i> # Assuming new tokens are of the format <s_i>
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])] 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 = (config.resolution, config.resolution)
original_size = (1024, 1024) original_size = (1024, 1024)
target_size = (config.resolution, config.resolution) target_size = (config.resolution, config.resolution)
crops_coords_top_left = (0,0)
crops_coords_top_left = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None: if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) 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) 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. 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) lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
else: else:
lora_prompt = prompt lora_prompt = prompt
@@ -299,6 +295,8 @@ def render_images(
is_lora, is_lora,
pretrained_model, pretrained_model,
lora_scale, lora_scale,
disable_ti=False,
prompt_modifier=None,
n_steps=25, n_steps=25,
n_imgs=4, n_imgs=4,
device="cuda:0", device="cuda:0",
@@ -328,6 +326,9 @@ def render_images(
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs) validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
validation_prompts_raw[0] = "<concept>" validation_prompts_raw[0] = "<concept>"
if prompt_modifier:
validation_prompts_raw = [prompt_modifier.format(prompt) for prompt in validation_prompts_raw]
if ( if (
checkpoint_folder is not None checkpoint_folder is not None
): # reload the entire pipeline from disk and load in the lora module ): # reload the entire pipeline from disk and load in the lora module
@@ -376,6 +377,7 @@ def render_images(
lora_scale, lora_scale,
guidance_scale=8, guidance_scale=8,
concept_mode=concept_mode, concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
) )
pipeline_args["prompt_embeds"] = c pipeline_args["prompt_embeds"] = c
@@ -415,6 +417,7 @@ def render_images_eval(
pretrained_model: dict, pretrained_model: dict,
trigger_text: str, trigger_text: str,
lora_scale=0.7, lora_scale=0.7,
disable_ti=False,
n_steps=25, n_steps=25,
n_imgs=4, n_imgs=4,
device="cuda:0", device="cuda:0",
@@ -469,6 +472,7 @@ def render_images_eval(
lora_scale, lora_scale,
guidance_scale=8, guidance_scale=8,
concept_mode=concept_mode, concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
) )
pipeline_args["prompt_embeds"] = c pipeline_args["prompt_embeds"] = c
+76 -1
View File
@@ -4,6 +4,81 @@ import matplotlib.pyplot as plt
import torch import torch
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
from trainer.inference import get_conditioning_signals 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): 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) self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
idx += 1 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 noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization 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 prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
+31 -51
View File
@@ -4,48 +4,30 @@ import subprocess
import torch import torch
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
############################################################################################################ def load_models(pretrained_model, device, weight_dtype = torch.float16):
SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
#SDXL_MODEL_CACHE = "./models/Juggernaut-X-RunDiffusion-NSFW.safetensors"
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-X-v10/resolve/main/Juggernaut-X-RunDiffusion-NSFW.safetensors"
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
#SD15_MODEL_CACHE = "./models/DreamShaper_6.31_BakedVae.safetensors"
#SD15_URL = "https://huggingface.co/Lykon/DreamShaper/resolve/main/DreamShaper_6.31_BakedVae.safetensors"
#SD15_MODEL_CACHE = "./models/photon_v1.safetensors"
#SD15_URL = "https://civitai.com/api/download/models/90072"
pretrained_models = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL, "version": "sd15"}
}
############################################################################################################
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
if not isinstance(pretrained_model, dict) or 'path' not in pretrained_model or 'version' not in pretrained_model:
raise ValueError("pretrained_model must be a dict with 'path' and 'version' keys")
# check if the model is already downloaded: # check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']): if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path']) download_weights(pretrained_model['url'], pretrained_model['path'])
print(f"Loading model weights from {pretrained_model['path']} with dtype: {weight_dtype}...") tokenizer_two, text_encoder_two = None, None
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
if pretrained_model['version'] == "sd15": try:
pipe = StableDiffusionPipeline.from_single_file( print("Loading as SDXL model...")
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
else:
pipe = StableDiffusionXLPipeline.from_single_file( pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True) 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) pipe = pipe.to(device, dtype=weight_dtype)
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config) noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
@@ -55,24 +37,11 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
text_encoder_one = pipe.text_encoder text_encoder_one = pipe.text_encoder
vae.requires_grad_(False) vae.requires_grad_(False)
if keep_vae_float32: vae.to(device, dtype=weight_dtype)
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 might not be for training..")
unet.to(device, dtype=weight_dtype) unet.to(device, dtype=weight_dtype)
text_encoder_one.requires_grad_(False) text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype) text_encoder_one.to(device, dtype=weight_dtype)
tokenizer_two = text_encoder_two = None
if pretrained_model['version'] == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
return ( return (
pipe, pipe,
tokenizer_one, tokenizer_one,
@@ -82,7 +51,7 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
text_encoder_two, text_encoder_two,
vae, vae,
unet, unet,
) ), sd_model_version
def download_weights(url, dest): def download_weights(url, dest):
start = time.time() start = time.time()
@@ -106,16 +75,27 @@ def download_weights(url, dest):
print(f"Downloading {url} took {time.time() - start} seconds") 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 trainable_params = 0
all_param = 0 all_param = 0
for name, param in model.named_parameters(): for name, param in model.named_parameters():
all_param += param.numel() all_param += param.numel()
if param.requires_grad and "token_embedding" not in name: if param.requires_grad and "token_embedding" not in name:
trainable_params += param.numel() 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 line_delimiter = "#" * 80
print(line_delimiter) print(line_delimiter)
print( 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) print(line_delimiter)
+42 -3
View File
@@ -13,9 +13,12 @@ def get_unet_optimizer(
): ):
## unet_trainable_params can be unet.parameters() or a list of lora params ## unet_trainable_params can be unet.parameters() or a list of lora params
# These learning rates will get overwritten in main.py:
if optimizer_name == "adamw": if optimizer_name == "adamw":
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0) optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
elif optimizer_name == "AdamW8bit":
import bitsandbytes as bnb
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
elif optimizer_name == "prodigy": elif optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT # Note: the specific settings of Prodigy seem to matter A LOT
optimizer_unet = prodigyopt.Prodigy( optimizer_unet = prodigyopt.Prodigy(
@@ -35,6 +38,39 @@ def get_unet_optimizer(
print(f"Created {optimizer_name} optimizer for unet!") print(f"Created {optimizer_name} optimizer for unet!")
return optimizer_unet return optimizer_unet
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
def is_belong_to_blocks(key, blocks):
try:
for g in blocks:
if g in key:
return True
return False
except Exception as e:
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
if use_blora:
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
try:
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
attns = [
attn_processor_name.rsplit(".", 1)[0]
for attn_processor_name, _ in unet.attn_processors.items()
if is_belong_to_blocks(attn_processor_name, blocks)
]
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
return target_modules
except Exception as e:
raise type(e)(
f"failed to get_target_modules, due to: {e}. "
f"Please check the modules specified in --lora_unet_blocks are correct"
)
def get_unet_lora_parameters( def get_unet_lora_parameters(
lora_rank, lora_rank,
lora_alpha_multiplier: float, lora_alpha_multiplier: float,
@@ -43,12 +79,15 @@ def get_unet_lora_parameters(
unet, unet,
pipe, pipe,
): ):
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
unet_lora_config = LoraConfig( unet_lora_config = LoraConfig(
r=lora_rank, r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier, lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian", init_lora_weights="gaussian",
target_modules=["to_k", "to_q", "to_v", "to_out.0", "conv2"], target_modules=target_modules,
#target_modules=["conv1", "conv2", "norm1", "norm2", "proj_in"], # TODO grid-search params for sd15
use_dora=use_dora, use_dora=use_dora,
) )
+135 -100
View File
@@ -1,7 +1,3 @@
# Have SwinIR upsample
# Have BLIP auto caption
# Have CLIPSeg auto mask concept
import gc import gc
import fnmatch import fnmatch
import mimetypes import mimetypes
@@ -25,6 +21,7 @@ import numpy as np
import pandas as pd import pandas as pd
import torch import torch
from tqdm import tqdm from tqdm import tqdm
from transformers import ( from transformers import (
BlipForConditionalGeneration, BlipForConditionalGeneration,
Blip2ForConditionalGeneration, Blip2ForConditionalGeneration,
@@ -38,13 +35,13 @@ from transformers import (
from trainer.utils.io import download_and_prep_training_data from trainer.utils.io import download_and_prep_training_data
from trainer.utils.utils import fix_prompt from trainer.utils.utils import fix_prompt
from trainer.config import model_paths
import re import re
import openai import openai
from openai import OpenAI from openai import OpenAI
from dotenv import load_dotenv from dotenv import load_dotenv
load_dotenv() load_dotenv()
try: try:
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY") OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
client = OpenAI(api_key=OPENAI_API_KEY) client = OpenAI(api_key=OPENAI_API_KEY)
@@ -54,11 +51,9 @@ except:
client = None client = None
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.") print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
MODEL_PATH = "./cache"
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...) # 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 MIN_GPT_PROMPTS = 3
MAX_GPT_PROMPTS = 50 MAX_GPT_PROMPTS = 80
def _find_files(pattern, dir="."): def _find_files(pattern, dir="."):
"""Return list of files matching pattern in a given directory, in absolute format. """Return list of files matching pattern in a given directory, in absolute format.
@@ -139,7 +134,7 @@ def swin_ir_sr(
""" """
model = Swin2SRForImageSuperResolution.from_pretrained( model = Swin2SRForImageSuperResolution.from_pretrained(
model_id, cache_dir=MODEL_PATH model_id, cache_dir = model_paths.get_path("SR")
).to(device) ).to(device)
processor = Swin2SRImageProcessor() processor = Swin2SRImageProcessor()
@@ -193,9 +188,9 @@ def clipseg_mask_generator(
model = None model = None
if any(target_prompts): if any(target_prompts):
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH) processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
model = CLIPSegForImageSegmentation.from_pretrained( model = CLIPSegForImageSegmentation.from_pretrained(
model_id, cache_dir=MODEL_PATH model_id, cache_dir = model_paths.get_path("CLIP")
).to(device) ).to(device)
masks = [] masks = []
@@ -236,7 +231,6 @@ def clipseg_mask_generator(
return masks return masks
import textwrap import textwrap
def cleanup_prompts_with_chatgpt( def cleanup_prompts_with_chatgpt(
prompts, prompts,
@@ -337,12 +331,12 @@ def extract_gpt_concept_description(gpt_completion, concept_mode):
return concept_name 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() text = text.strip()
gpt_cleanup_worked = False gpt_cleanup_worked = False
gpt_concept_description = None 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 retry_count = 0
while retry_count < 5: while retry_count < 5:
try: try:
@@ -408,14 +402,14 @@ def blip_caption_dataset(
device=torch.device("cuda" if torch.cuda.is_available() else "cpu") device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "blip2" in model_id: if "blip2" in model_id:
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH) processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = Blip2ForConditionalGeneration.from_pretrained( model = Blip2ForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16 model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device) ).to(device)
else: else:
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH) processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = BlipForConditionalGeneration.from_pretrained( model = BlipForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16 model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device) ).to(device)
for i, image in enumerate(tqdm(images)): for i, image in enumerate(tqdm(images)):
@@ -424,6 +418,7 @@ def blip_caption_dataset(
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65) 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) captions[i] = processor.decode(out[0], skip_special_tokens=True)
model.to("cpu")
del model del model
gc.collect() gc.collect()
torch.cuda.empty_cache() torch.cuda.empty_cache()
@@ -445,51 +440,6 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
os.remove(output_path) os.remove(output_path)
return base64_image 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-4-turbo",
"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( def gpt4_v_caption_dataset(
images, captions, images, captions,
batch_size=4, batch_size=4,
@@ -500,6 +450,7 @@ def gpt4_v_caption_dataset(
return captions 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 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 = { headers = {
"Content-Type": "application/json", "Content-Type": "application/json",
@@ -510,7 +461,7 @@ def gpt4_v_caption_dataset(
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512)) base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
payload = { payload = {
"model": "gpt-4-turbo", "model": "gpt-4o",
"messages": [ "messages": [
{ {
"role": "user", "role": "user",
@@ -548,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() @torch.no_grad()
def caption_dataset( def caption_dataset(
images: List[Image.Image], images: List[Image.Image],
captions: List[str], captions: List[str],
caption_model: Literal[str] = "blip" caption_model: Literal["blip", "gpt4-v", "florence"] = "blip"
) -> List[str]: ) -> 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: if "blip" in caption_model:
captions = blip_caption_dataset(images, captions) captions = blip_caption_dataset(images, captions)
elif "gpt4-v" in caption_model: elif "gpt4-v" in caption_model:
captions = gpt4_v_caption_dataset(images, captions) 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 return captions
@@ -629,20 +647,44 @@ def random_crop(image, scale=(0.85, 0.95)):
return image.crop((left, top, left + new_width, top + new_height)) return image.crop((left, top, left + new_width, top + new_height))
def gaussian_blur(image): def gaussian_blur(image, radius = 1.0):
return image.filter(ImageFilter.GaussianBlur(radius=1)) return image.filter(ImageFilter.GaussianBlur(radius=radius))
def augment_image(image): def augment_image(image):
image = hue_augmentation(image) image = hue_augmentation(image)
image = color_jitter(image) image = color_jitter(image)
image = random_crop(image) image = random_crop(image)
if random.random() < 0.5: if random.random() < 0.5:
image = gaussian_blur(image) image = gaussian_blur(image, radius = random.uniform(0.0, 1.0))
return image return image
def round_to_nearest_multiple(x, multiple): def round_to_nearest_multiple(x, multiple):
return int(float(multiple) * round(float(x) / float(multiple))) return int(float(multiple) * round(float(x) / float(multiple)))
'''
For Stable Diffusion 1.5, outputs are optimised around 512x512 pixels. Many common fine-tuned versions of SD1.5 are optimised around 768x768. The best resolutions for common aspect ratios are typically:
1:1 (square): 512x512, 768x768
3:2 (landscape): 768x512
2:3 (portrait): 512x768
4:3 (landscape): 768x576
3:4 (portrait): 576x768
16:9 (widescreen): 912x512
9:16 (tall): 512x912
For SDXL, outputs are optimised around 1024x1024 pixels. The best resolutions for common aspect ratios are typically:
stable-diffusion-xl-1024-v0-9 supports generating images at the following dimensions:
1024 x 1024
1152 x 896
896 x 1152
1216 x 832
832 x 1216
1344 x 768
768 x 1344
1536 x 640
640 x 1536
'''
def calculate_new_dimensions(target_size, target_aspect_ratio): def calculate_new_dimensions(target_size, target_aspect_ratio):
""" """
Calculate the new width and height given a target size and aspect ratio. Calculate the new width and height given a target size and aspect ratio.
@@ -661,8 +703,6 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
return [new_width, new_height] return [new_width, new_height]
def load_and_save_masks_and_captions( def load_and_save_masks_and_captions(
config, config,
concept_mode: str, concept_mode: str,
@@ -703,9 +743,10 @@ def load_and_save_masks_and_captions(
n_length = len(files) n_length = len(files)
files = sorted(files)[:n_length] files = sorted(files)[:n_length]
images, captions = [], [] images, captions, img_paths = [], [], []
for file in files: for file in files:
images.append(load_image_with_orientation(file)) images.append(load_image_with_orientation(file))
img_paths.append(file)
caption_file = os.path.splitext(file)[0] + ".txt" caption_file = os.path.splitext(file)[0] + ".txt"
if os.path.exists(caption_file) and use_dataset_captions: if os.path.exists(caption_file) and use_dataset_captions:
with open(caption_file, "r") as f: with open(caption_file, "r") as f:
@@ -746,41 +787,39 @@ def load_and_save_masks_and_captions(
upscale_margin = 0.75 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))) 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})") 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] images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
captions = captions + captions 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) print(f"Generating {len(images)} captions using {caption_model} in {concept_mode} mode...")
captions.extend(aug_caps) 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: # 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): if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
images = images[:MAX_GPT_PROMPTS-1] images = images[:MAX_GPT_PROMPTS-1]
captions = captions[: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: # Cleanup prompts using chatgpt:
captions = [fix_prompt(caption) for caption in captions] captions = [fix_prompt(caption) for caption in captions]
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed) 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, 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 = [],[] 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: while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
print(f"Adding augmented version of each training img...") print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images]) aug_imgs.extend([augment_image(image) for image in images])
@@ -789,19 +828,18 @@ def load_and_save_masks_and_captions(
images.extend(aug_imgs) images.extend(aug_imgs)
captions.extend(aug_caps) captions.extend(aug_caps)
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")): 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}") print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
mask_target_prompts = gpt_concept_description mask_target_prompts = gpt_concept_description
if mask_target_prompts is None or config.concept_mode == "style": if mask_target_prompts is None or config.concept_mode == "style":
print("Disabling CLIP-segmentation")
mask_target_prompts = "" mask_target_prompts = ""
temp = 999 temp = 999
else: else:
temp = config.clipseg_temperature temp = config.clipseg_temperature
print(f"Generating {len(images)} masks...") print(f"Generating {len(images)} masks...")
# Make sure we have a bias for the background pixels to never 100% ignore them # Make sure we have a bias for the background pixels to never 100% ignore them
background_bias = 0.05 background_bias = 0.05
if not use_face_detection_instead: if not use_face_detection_instead:
@@ -855,15 +893,10 @@ def load_and_save_masks_and_captions(
os.makedirs(output_dir, exist_ok=True) os.makedirs(output_dir, exist_ok=True)
# Make sure we've correctly inserted the TOK into every caption: if config.disable_ti:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
for caption in captions:
print(caption)
if config.remove_ti_token_from_prompts:
print('------------------ WARNING -------------------') print('------------------ WARNING -------------------')
print("Removing 'TOK, ' from captions...") print("Removing 'TOK, ' from captions...")
print("This will completely break textual_inversion!!") print("This will completely disable textual_inversion!!")
print('------------------ WARNING -------------------') print('------------------ WARNING -------------------')
if gpt_concept_description: if gpt_concept_description:
replace_str = gpt_concept_description replace_str = gpt_concept_description
@@ -871,6 +904,8 @@ def load_and_save_masks_and_captions(
replace_str = "" replace_str = ""
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions] captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
captions = [caption.replace("TOK", replace_str) for caption in captions] captions = [caption.replace("TOK", replace_str) for caption in captions]
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
# iterate through the images, masks, and captions and add a row to the dataframe for each # iterate through the images, masks, and captions and add a row to the dataframe for each
print("Saving final training dataset...") print("Saving final training dataset...")
+364
View File
@@ -0,0 +1,364 @@
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 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):
module_names = []
"""
this function assumes that there are fewer than 50 down blocks, attention modules and transformer blocks
if you're not sure, feel free to set it to an arbitrarily large number
don't worry, it won't slow anything down.
"""
for block_type in ["down_blocks", "up_blocks"]:
for down_block_index in range(50):
for attentions_index in range(50):
for transformer_blocks_index in range(50):
example_module_name = f"{block_type}.{down_block_index}.attentions.{attentions_index}.transformer_blocks.{transformer_blocks_index}.attn2.processor"
try:
module = get_module_by_name(module=unet, name = example_module_name)
assert isinstance(module, AttnProcessor2_0), f"Expected module to be an instance of AttnProcessor2_0 but found it to be: {type(module)}"
# print(f"Found: {example_module_name}")
module_names.append(example_module_name)
except AttributeError:
# print(f"Ignored name: {example_module_name}\nsince it does not exist")
pass
print(f"Found: {len(module_names)} modules")
return module_names
class DAAMLossAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(self, name: str):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
self.name = name
self.cross_attention_scores = None
self.reduce_op = Reduce(
"batch heads img text -> batch img text",
reduction="sum"
)
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
temb: Optional[torch.Tensor] = None,
*args,
**kwargs,
) -> torch.Tensor:
if len(args) > 0 or kwargs.get("scale", None) is not None:
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
deprecate("scale", "1.0.0", deprecation_message)
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
"""
Mayukh's experiment
"""
mayukh_experiment = False
if encoder_hidden_states is not None:
"""
this triggers cross attn
"""
mayukh_experiment = True
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
if mayukh_experiment:
# Calculate QK^T
qk_t = torch.matmul(query, key.transpose(-2, -1))
# Calculate attention scores (scaled QK^T)
d_k = query.size(-1) # Assuming the last dimension is the embedding dimension
attention_scores = qk_t / math.sqrt(d_k)
attention_scores = self.reduce_op(
attention_scores,
)
self.cross_attention_scores = attention_scores
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class DAAMLoss:
def __init__(self, attention_processors: list[DAAMLossAttnProcessor2_0]):
self.attention_processors = attention_processors
self.layer_names = [
x.name for x in attention_processors
]
def 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 = {}
for p in self.attention_processors:
cross_attention_scores[
p.name
] = p.cross_attention_scores
return cross_attention_scores
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]
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 = height,
width = width
)
return heatmap
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,
img_ratio = img_ratio
)
all_heatmaps.append(heatmap)
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, heatmap_height, heatmap_width) using F.interpolate
resized_heatmaps = [
F.interpolate(input = x.unsqueeze(1), size = (heatmap_height, heatmap_width)).squeeze(1)
for x in all_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."""
if name == "":
return module
names = name.split(sep=".")
return reduce(getattr, names, module)
def init_daam_loss(pipeline):
## find out where the attention processor thingies are
module_names = find_attnprocessor2_0(
unet = pipeline.unet
)
all_daam_attention_processors = []
# override the attention processor thingies
for name in module_names:
# Get parent module and attribute name
parent_name = ".".join(name.split(".")[:-1])
attr_name = name.split(".")[-1]
# Get the parent module
parent_module = get_module_by_name(module=pipeline.unet, name=parent_name)
daam_attention_processor = DAAMLossAttnProcessor2_0(name=name)
all_daam_attention_processors.append(daam_attention_processor)
# Set the attribute
setattr(parent_module, attr_name, daam_attention_processor)
# Verify the replacement
current_module = get_module_by_name(module=pipeline.unet, name=name)
assert isinstance(current_module, DAAMLossAttnProcessor2_0)
daam_loss = DAAMLoss(
attention_processors=all_daam_attention_processors
)
return pipeline, daam_loss
+12
View File
@@ -290,6 +290,18 @@ def load_image_with_orientation(path, mode="RGB"):
elif orientation == 8: elif orientation == 8:
image = image.rotate(90, expand=True) 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) return image.convert(mode)
+13 -4
View File
@@ -100,15 +100,17 @@ def print_system_info():
# Print disk space information # Print disk space information
disk_usage = psutil.disk_usage('/') disk_usage = psutil.disk_usage('/')
free_disk = disk_usage.free // (1024 * 1024) total_disk = disk_usage.total // (1024 * 1024)
used_disk = disk_usage.used // (1024 * 1024)
percent_disk_used = disk_usage.percent percent_disk_used = disk_usage.percent
print(f"Free disk space: {free_disk} MB with {percent_disk_used}% used") print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
# Print RAM information # Print RAM information
virtual_mem = psutil.virtual_memory() virtual_mem = psutil.virtual_memory()
total_ram = virtual_mem.total // (1024 * 1024)
current_ram = virtual_mem.used // (1024 * 1024) current_ram = virtual_mem.used // (1024 * 1024)
percent_ram_used = virtual_mem.percent percent_ram_used = virtual_mem.percent
print(f"Current used RAM: {current_ram} MB with {percent_ram_used}% used") print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
except Exception as e: except Exception as e:
print(f'Error in gathering system info: {str(e)}') print(f'Error in gathering system info: {str(e)}')
@@ -122,6 +124,13 @@ def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1
# Flatten and concatenate all parameters into a single tensor # Flatten and concatenate all parameters into a single tensor
all_params = torch.cat([p.data.view(-1) for p in parameters]) all_params = torch.cat([p.data.view(-1) for p in parameters])
# count number of parameters:
n_params = len(all_params)
if n_params == 0 or n_params > 1e9:
return
norm = torch.norm(all_params) norm = torch.norm(all_params)
# Convert to CPU for plotting # Convert to CPU for plotting
@@ -228,7 +237,7 @@ def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dic
from scipy.signal import savgol_filter from scipy.signal import savgol_filter
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'): 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'] values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
plot_smoothed = ['img_loss'] plot_smoothed = ['img_loss']
+2 -2
View File
@@ -2,9 +2,9 @@
val_prompts = {} val_prompts = {}
val_prompts['style'] = [ val_prompts['style'] = [
'a beautiful mountainous landscape, boulders, fresh water stream, setting sun', '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', '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 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 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', 'a stunning image of an aston martin sportscar',
-31
View File
@@ -1,31 +0,0 @@
{
"output_dir": "lora_models/xander_sd15_final",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 16,
"unet_lr": 0.001,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": false,
"caption_model": "gpt4-v",
"debug": true
}
-25
View File
@@ -1,25 +0,0 @@
{
"output_dir": "lora_models/object",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/DOV/lizzo/full body",
"concept_mode": "object",
"seed": 1,
"resolution": 640,
"train_batch_size": 4,
"n_sample_imgs": 4,
"max_train_steps": 420,
"token_warmup_steps": 0,
"checkpointing_steps": 60,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"lora_rank": 12,
"use_dora": false,
"caption_model": "gpt4-v",
"debug": true
}
-29
View File
@@ -1,29 +0,0 @@
{
"output_dir": "lora_models/does_best",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip",
"concept_mode": "style",
"seed": 1,
"resolution": 640,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
-29
View File
@@ -1,29 +0,0 @@
{
"output_dir": "lora_models/Journey",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/journey",
"concept_mode": "style",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 1000,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"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,
"prodigy_d_coef": 1.0,
"unet_prodigy_growth_factor": 1.05,
"lora_rank": 16,
"use_dora": true,
"caption_model": "gpt4-v",
"debug": true
}