67 Commits
Author SHA1 Message Date
mayukhdeb 6e4c86ea0c keep a copy 2024-07-22 11:07:14 -07:00
mayukhdeb 7e30d33dc2 find image filenames recursively in folder + train on big style dataset 2024-07-21 03:56:42 -07:00
mayukhdeb fbd050cb0c cleaner preprocess fn 2024-07-21 03:25:07 -07:00
mayukhdeb 70c8e88195 migrate to huggingface script 2024-07-21 02:05:30 -07:00
mayukhdeb 2f6f0cacd2 ignore stuff 2024-07-21 02:03:42 -07:00
mayukhdeb e8932f30ef start with a different embed string 2024-07-18 09:57:00 -07:00
mayukhdeb b7dd176536 add adamw8bit 2024-07-15 07:48:07 -07:00
mayukhdeb 5a72aca596 useful wandb name 2024-07-15 07:47:34 -07:00
mayukhdeb 85264bd004 wandb 2024-07-15 07:38:27 -07:00
mayukhdeb d46976f3dd sweep params set 2024-07-15 07:33:45 -07:00
mayukhdeb 8ed5996c1f train and inference on same device 2024-07-15 07:32:27 -07:00
mayukhdeb 8247b4b9dd fix OOM during inference (vae.decode) 2024-07-15 07:20:13 -07:00
mayukhdeb dc28fcb988 inference on same device 2024-07-15 03:28:31 -07:00
mayukhdeb 3c0a3be55f save different sh files for each gpu 2024-07-15 03:28:20 -07:00
mayukhdeb 274ec5c90b disable textual inversion if config.ti_lr is None 2024-07-11 12:46:52 -07:00
mayukhdeb fc54948eff more tweaks 2024-07-11 12:39:30 -07:00
mayukhdeb 32a7ae248d prompts for sweep 2024-07-11 12:34:42 -07:00
mayukhdeb 3cbabd2634 sweep dry runs 2024-07-11 12:26:07 -07:00
mayukhdeb 1cdacbc63a remove old todos 2024-07-11 11:16:22 -07:00
mayukhdeb 4668f755ec switch to adamw_8bit for sd3 transformer lora params 2024-07-11 00:15:09 -07:00
mayukhdeb dd093616b5 fix inference prompts bug 2024-07-09 09:09:56 -07:00
mayukhdeb 2c824d6807 more progress 2024-07-09 05:35:21 -07:00
mayukhdeb 022d51f53a inference fixed 2024-07-09 00:25:02 -07:00
mayukhdeb ecacf815ec re-impl textual inversion for first 2 text encoders 2024-07-08 12:45:55 -07:00
mayukhdeb 90d2269572 testing textual inversion training with frozen transformer 2024-07-04 07:40:41 -07:00
mayukhdeb 4955261cae new checkpoint + deterministic inference 2024-07-03 04:28:06 -07:00
mayukhdeb e84932af29 some small changes to text with main_sd3.py 2024-07-03 04:06:21 -07:00
mayukhdeb 115da83e2d cleaner output dir with checkpoints and generated samples in one folder 2024-07-03 03:43:30 -07:00
mayukhdeb 98b78cd6d6 cleanup + save training samples in output dir 2024-07-03 03:11:56 -07:00
mayukhdeb 7c5e0949ba keep changes 2024-07-02 05:19:34 -07:00
mayukhdeb ba2b7532ad more tweaks 2024-07-02 05:16:22 -07:00
mayukhdeb 2c0939a733 run inference less often 2024-07-02 04:25:26 -07:00
mayukhdeb 3b41479aa1 bfloat16 training + inference 2024-07-02 04:12:23 -07:00
mayukhdeb 6e589c76a8 impement some upstream changes and save a sample every 10 train steps 2024-07-02 01:50:14 -07:00
mayukhdeb 95f78d7b91 completely comement out TI for now 2024-07-01 01:56:49 -07:00
mayukhdeb 3c0ebd7d70 apply mask to loss + some hardcoding for banny debugging 2024-06-28 04:38:43 -07:00
mayukhdeb 6350b5344b more progress 2024-06-28 03:22:49 -07:00
mayukhdeb 2c4bd43044 temporarily remove ti grad norms 2024-06-28 03:12:16 -07:00
mayukhdeb b850a1615d watch grad norms 2024-06-28 03:04:40 -07:00
mayukhdeb 6d6dd98bfa dynamic ti lr 2024-06-28 02:03:34 -07:00
mayukhdeb 0444e35729 clip grad norms 2024-06-26 03:01:34 -07:00
mayukhdeb 0748e1d14e better prompt 2024-06-22 02:32:19 -07:00
mayukhdeb db7507c849 handle T5EncoderModel 2024-06-22 01:38:48 -07:00
mayukhdeb b379a28715 update todos 2024-06-22 01:38:12 -07:00
mayukhdeb 54ff8c4977 sd3 concept inference 2024-06-22 01:25:16 -07:00
mayukhdeb 64f28c8589 save TI embeds and lora adapters 2024-06-22 01:24:30 -07:00
mayukhdeb b12ce26fc9 update command 2024-06-20 03:52:30 -07:00
mayukhdeb b3da65dd39 ignore wandb stuff 2024-06-20 03:46:26 -07:00
mayukhdeb f64da5d4f1 update todo 2024-06-20 03:44:42 -07:00
mayukhdeb 5cc09092a4 tweak param 2024-06-20 03:43:44 -07:00
mayukhdeb 3d3d0ee4cb smash more todos 2024-06-20 03:42:32 -07:00
mayukhdeb 4b1efce0e3 compute loss and update weights 2024-06-20 03:28:49 -07:00
mayukhdeb d115c52274 do just forward passes 2024-06-20 03:20:41 -07:00
mayukhdeb 22127c9917 typo 2024-06-20 01:41:28 -07:00
mayukhdeb 4844845d5c more progress 2024-06-20 01:40:59 -07:00
mayukhdeb 3b15250bbc init train dataloader + update todos for training 2024-06-20 00:54:41 -07:00
mayukhdeb fd448e5437 accomodate T5EncoderModel 2024-06-20 00:54:21 -07:00
mayukhdeb a41cce1486 small cleanup 2024-06-20 00:27:51 -07:00
mayukhdeb 12b5960c36 progress bar for latent caching 2024-06-20 00:27:21 -07:00
mayukhdeb 2c2bfcda13 init PreprocessedDataset 2024-06-20 00:27:01 -07:00
mayukhdeb cacd6f2201 count trainable params from model 2024-06-19 23:58:26 -07:00
mayukhdeb 67024b4de6 full or lora finetuning of sd3 transformer 2024-06-19 23:58:16 -07:00
mayukhdeb 6d0d96ec79 more progress on todos 2024-06-19 23:37:31 -07:00
mayukhdeb fb2449a60d init textual inversion token embeds 2024-06-17 07:31:54 -07:00
mayukhdeb 0496ccdbef handle sd3 t5 text encoder 2024-06-17 07:31:33 -07:00
mayukhdeb 044aed9f03 sd3 train script wip 2024-06-17 06:42:16 -07:00
mayukhdeb bdda796c56 ignore notebook checkpoint 2024-06-17 05:21:54 -07:00
53 changed files with 3654 additions and 2827 deletions
+8 -1
View File
@@ -22,4 +22,11 @@ datasets
rendered_images
# 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
@@ -1,26 +0,0 @@
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 }}
+9 -5
View File
@@ -1,25 +1,29 @@
data/
sd3_sweep_vis/
sd3_sweep_commands/
.ipynb_checkpoints/
cache
__pycache__
.ipynb_checkpoints/
models
lora_models*
eden_lora_training_runs/
datasets
*.tar
.env
.cog
.huggingface
train.py
rendered_images*
gridsearch*
aesthetic_score_best_model.pth
# experiment folders:
scripts/plots
conditioning_spaces/
training_args_x_*.json
xander_configs/
debug/*
wandb/
sd3_sweep_outputs/
sd3_face_sweep_configs/
@@ -1,782 +0,0 @@
{
"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
@@ -1,312 +0,0 @@
{
"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
@@ -1,43 +0,0 @@
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.
+50 -30
View File
@@ -1,11 +1,10 @@
# Trainer
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/).
This trainer was developed by the [**Eden** team](https://eden.art/)
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
The outputs of this trainer are fully compatible with ComfyUI and AUTO111, 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).
The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
<p align="center">
<strong>Training images:</strong><br>
@@ -17,26 +16,11 @@ A full guide on training can be found in [**our docs**](https://docs.eden.art/do
</p>
### 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:
The trainer supports 3 default modes:
- **style**: used for learning the aesthetic style of a collection of images.
- **face**: used for learning a specific face (can be human, character, ...).
- **object**: will learn a specific object or thing featured in the training images.
<p align="center">
<strong>Style training example:</strong><br>
<img src="assets/style_training_example.jpg" alt="Image 1" style="width:80%;"/>
</p>
## Setup
Install all dependencies using
@@ -45,7 +29,7 @@ Install all dependencies using
then you can simply run:
`python main.py train_configs/training_args.json`
`python main.py -c training_args.json`
to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job.
@@ -60,25 +44,61 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
sudo chmod +x /usr/local/bin/cog
```
2. Build the image with `cog build`
3. Run a training run with `sh cog_test_train.sh`
4. You can also go into the container with `cog run /bin/bash`
2. Build the image with `sudo cog build`
3. Run a training run with `sudo sh cog_test_train.sh`
## Automatic Checkpoint Evaluation
This script uses CLIP img/txt similarity scores to evaluate how good the LoRa is vs how overfit.
Download the aesthetic predictor model checkpoint first from google drive. This should give you a file named: `aesthetic_score_best_model.pth` (99.2 MB)
```bash
gdown 1thEIlXVc8lkULVUBY9Ab45tsOERxkjxns
```
Once the model is downloaded, you can run the eval script with the following CLI args:
- `output_folder`: this is where the outputs of the model get saved as jpeg files
- `lora_path`: path to your LoRA checkpoint (make sure you edit `path_to_your_model_checkpoints` to point to the correct folder. It generally ends with something like `checkpoint-600` where `600` was the training step)
- `output_json`: save all scores in this json file
- `config_filename`: config file used for training
```bash
python3 evaluate.py \
--output_folder eval_images \
--lora_path path_to_your_model_checkpoint \
--output_json eval_results.json \
--config_filename training_args.json
```
## Full unet finetuning
When running this trainer in native python, you can also perform full unet finetuning using something like (adjust to your needs)
`python main.py train_configs/full_finetuning_example.json`
## TODO's
Bugs:
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
- 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)
- 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:
- integrate Flux / SD3
- Add multi-concept training (multiple things represented by multiple tokens, trained into a single LoRa)
- add stronger token regularization (eg CelebBasis spanning basis)
- add stronger token regularization (eg CelebBasis spanning basis):
- Add multi-token training
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- implement prompt-aligned: https://prompt-aligned.github.io/
Tuning Experiments:
- try-out conditioning noise injection during training to increase robustness
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
but it then recovers. Can we avoid this collapse? Is the learning rate too high?
- offset noise
- AB test Dora vs Lora
Binary file not shown.

Before

Width:  |  Height:  |  Size: 430 KiB

+7 -4
View File
@@ -3,14 +3,17 @@
build:
gpu: true
cuda: "12.1"
python_version: "3.11"
cuda: "11.8"
python_version: "3.9"
system_packages:
- "ffmpeg"
- "libgl1-mesa-glx"
- "libegl1-mesa-dev"
- "libsm6"
- "libxext6"
python_requirements: requirements.txt
run:
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
- 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
predict: "predict.py:Predictor"
image: "r8.im/edenartlab/sdxl-lora-trainer"
+6 -7
View File
@@ -1,13 +1,12 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=3"
GPU_ID="device=2"
cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
-i concept_mode="style" \
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip" \
-i concept_mode="face" \
-i sd_model_version="sdxl" \
-i max_train_steps="300" \
-i sample_imgs_lora_scale="0.7" \
-i n_sample_imgs="6" \
-i debug="True" \
-i max_train_steps="360" \
-i caption_model="blip" \
-i debug="False" \
-i seed="0"
+533
View File
@@ -0,0 +1,533 @@
{
"last_node_id": 12,
"last_link_id": 23,
"nodes": [
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 16
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
6
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"text, watermark"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1209,
188
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 8
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
9
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 12,
"type": "Reroute",
"pos": [
220,
14
],
"size": [
75,
26
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 22
}
],
"outputs": [
{
"name": "",
"type": "MODEL",
"links": [
19
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 3,
"type": "KSampler",
"pos": [
863,
186
],
"size": {
"0": 315,
"1": 262
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 19
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 4
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 6
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
7
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
1,
"fixed",
25,
8,
"euler",
"normal",
1
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
473,
609
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
768,
768,
1
]
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-467,
120
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
10
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
12
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
8
],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"juggernaut_reborn.safetensors"
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
415,
186
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 15
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
4
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"a photo of embedding:xander_sd15_embedding on the beach "
]
},
{
"id": 11,
"type": "Reroute",
"pos": [
220,
49
],
"size": [
75,
26
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 23
}
],
"outputs": [
{
"name": "",
"type": "CLIP",
"links": [
15,
16
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 9,
"type": "SaveImage",
"pos": [
347,
-270
],
"size": {
"0": 210,
"1": 270
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 9
}
],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 10,
"type": "LoraLoader",
"pos": [
-72,
-229
],
"size": {
"0": 254.95774841308594,
"1": 127.86701202392578
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 10
},
{
"name": "clip",
"type": "CLIP",
"link": 12
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
22
],
"shape": 3,
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
23
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LoraLoader"
},
"widgets_values": [
"xander_sd15_lora.safetensors",
0.6,
0.6
]
}
],
"links": [
[
2,
5,
0,
3,
3,
"LATENT"
],
[
4,
6,
0,
3,
1,
"CONDITIONING"
],
[
6,
7,
0,
3,
2,
"CONDITIONING"
],
[
7,
3,
0,
8,
0,
"LATENT"
],
[
8,
4,
2,
8,
1,
"VAE"
],
[
9,
8,
0,
9,
0,
"IMAGE"
],
[
10,
4,
0,
10,
0,
"MODEL"
],
[
12,
4,
1,
10,
1,
"CLIP"
],
[
15,
11,
0,
6,
0,
"CLIP"
],
[
16,
11,
0,
7,
0,
"CLIP"
],
[
19,
12,
0,
3,
0,
"MODEL"
],
[
22,
10,
0,
12,
0,
"*"
],
[
23,
10,
1,
11,
0,
"*"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": {
"0": 513.8734070325743,
"1": 351.4824273966635
}
}
},
"version": 0.4
}
+28
View File
@@ -0,0 +1,28 @@
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)}')
+179
View File
@@ -0,0 +1,179 @@
from trainer.utils.json_stuff import save_as_json
import itertools
import copy
import os
import random
random.seed(0)
GPU_IDS = [1,2,3]
wandb_log = True
def divide_list(lst, n):
"""
Divide a list into N equal parts.
Parameters:
lst (list): The list to be divided.
n (int): The number of parts to divide the list into.
Returns:
list of lists: A list containing N sublists, each of which is a part of the original list.
"""
if n <= 0:
raise ValueError("Number of parts must be greater than 0.")
if n > len(lst):
raise ValueError("Number of parts cannot be greater than the length of the list.")
# Calculate the size of each part
k, m = divmod(len(lst), n)
# Create the divided parts
return [lst[i * k + min(i, m):(i + 1) * k + min(i + 1, m)] for i in range(n)]
def generate_sh_file(commands, filename="script.sh"):
"""
Generates a .sh file with each command from the list written on a new line.
:param commands: List of commands to be written to the .sh file.
:param filename: Name of the .sh file to be created. Default is 'script.sh'.
"""
with open(filename, 'w') as file:
for command in commands:
file.write(command + '\n')
print(f"Saved: {filename}")
run_commands_dir = f"./sd3_sweep_commands"
os.system(
f"rm -rf {run_commands_dir} && mkdir -p {run_commands_dir}"
)
config_folder = "./sd3_face_sweep_configs"
os.system(f"rm -rf {config_folder}")
os.system(f"mkdir -p {config_folder}")
sweep_params = {
"unet_learning_rate": [
5e-5,
1e-4,
3e-4,
7e-4,
1e-3,
2e-3,
],
"train_batch_size": [
2,
4,
8,
16
],
"lora_rank": [
2,
4,
6,
8,
],
"ti_lr": [1e-3, None],
"unet_optimizer_type": [
"adamw",
"adamw_8bit",
"prodigy"
],
}
num_total_runs = 1
for key in sweep_params:
num_total_runs *= len(sweep_params[key])
print(f"Num total runs: {num_total_runs}")
default_config = {
"output_dir": "lora_models/sweep",
"sd_model_version": "sd3",
"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": 2,
"n_sample_imgs": 6,
"max_train_steps": 1000,
"token_warmup_steps": 200,
"checkpointing_steps": 1000, ## no need to save any checkpoints
"gradient_accumulation_steps": 2,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": False,
"text_encoder_lora_optimizer": None,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 16,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": False,
"caption_model": "blip",
"debug": True,
}
keys, values = zip(*sweep_params.items())
combinations = [dict(zip(keys, combination)) for combination in itertools.product(*values)]
all_config_paths = []
for index, c in enumerate(combinations):
config = copy.deepcopy(default_config)
filename = f"{index}"
# override default values with sweep params
for key in c:
"""
instead of editing the train batch size, we simply change the gradient
accumulation value. Which has the same effect.
We will also
"""
if key == "train_batch_size":
config["gradient_accumulation_steps"] = c[key] / config["train_batch_size"]
config["max_train_steps"] = config["max_train_steps"] * config["gradient_accumulation_steps"]
config["checkpointing_steps"] = config["checkpointing_steps"] * config["gradient_accumulation_steps"]
else:
config[key] = c[key]
# print(f"{index} - Setting {key} to {c[key]}")
filename += f"_{key}_{c[key]}"
config_path = os.path.join(
config_folder,
f"{filename}.json"
)
save_as_json(
dictionary_or_list=config,
filename = config_path
)
all_config_paths.append(config_path)
print(f"Saved: {config_path}")
print(f"Total: {index+1} configs")
all_commands = []
for c in all_config_paths:
command = f"python3 main_sd3.py {c}"
if wandb_log:
command = command + " --wandb-log"
all_commands.append(command)
random.shuffle(all_commands)
all_commands_split_by_gpu = divide_list(
lst = all_commands,
n = len(GPU_IDS)
)
for index, gpu_id in enumerate(GPU_IDS):
commands_on_single_gpu = [
f"CUDA_VISIBLE_DEVICES={gpu_id} {x}" for x in all_commands_split_by_gpu[index]
]
generate_sh_file(
commands = commands_on_single_gpu,
filename = os.path.join(
run_commands_dir,
f"run_on_gpu_{gpu_id}.sh"
)
)
+116 -137
View File
@@ -1,3 +1,4 @@
import fnmatch
import math
import os
import time
@@ -11,18 +12,20 @@ import torch
import torch.utils.checkpoint
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.checkpoint import save_checkpoint
from trainer.embedding_handler import TokenEmbeddingsHandler
from trainer.dataset import PreprocessedDataset
from trainer.config import TrainingConfig
from trainer.models import print_trainable_parameters, load_models
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer, compute_token_attention_loss
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer
from trainer.inference import render_images, get_conditioning_signals
from trainer.preprocess import preprocess
from trainer.utils.io import make_validation_img_grid
from trainer.ti_cross_attn_loss import init_daam_loss, plot_token_attention_loss
from trainer.optimizer import (
OptimizerCollection,
get_optimizer_and_peft_models_text_encoder_lora,
@@ -31,43 +34,10 @@ from trainer.optimizer import (
get_unet_optimizer
)
def train(config: TrainingConfig):
def train(
config: TrainingConfig,
):
seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
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,
@@ -88,6 +58,19 @@ def train(config: TrainingConfig):
if config.allow_tf32:
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.
embedding_handler = TokenEmbeddingsHandler(
text_encoders = [text_encoder_one, text_encoder_two],
@@ -130,26 +113,22 @@ def train(config: TrainingConfig):
embedding_handler.make_embeddings_trainable()
if not config.disable_ti:
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
else:
optimizer_ti = None
textual_inversion_params = None
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
if not config.is_lora: # This code pathway has not been tested in a long while
print(f"Doing full fine-tuning on the U-Net")
unet.requires_grad_(True)
unet_lora_parameters = None
optimizer_text_encoder_lora = None
unet_trainable_params = unet.parameters()
else:
# Do lora-training instead.
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
# target_blocks=["block"] for original IP-Adapter
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
@@ -163,18 +142,15 @@ def train(config: TrainingConfig):
pipe=pipe
)
if config.unet_lr > 0.0:
optimizer_unet = get_unet_optimizer(
prodigy_d_coef=config.prodigy_d_coef,
prodigy_growth_factor=config.unet_prodigy_growth_factor,
lora_weight_decay=config.lora_weight_decay,
use_dora=config.use_dora,
unet_trainable_params=unet_trainable_params,
optimizer_name=config.unet_optimizer_type
)
else:
optimizer_unet = None
optimizer_unet = get_unet_optimizer(
prodigy_d_coef=config.prodigy_d_coef,
prodigy_growth_factor=config.unet_prodigy_growth_factor,
lora_weight_decay=config.lora_weight_decay,
use_dora=config.use_dora,
unet_trainable_params=unet_trainable_params,
optimizer_name=config.unet_optimizer_type
)
print_trainable_parameters(unet, model_name = 'unet')
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
@@ -185,26 +161,31 @@ def train(config: TrainingConfig):
pipe,
vae.float(),
size = config.train_img_size,
do_cache=config.do_cache,
substitute_caption_map=config.token_dict,
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
train_batch_size=config.train_batch_size
)
print("Final training captions:")
print(train_dataset.captions[:40])
# offload the vae to cpu and release memory:
# offload the vae to cpu:
vae = vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=config.train_batch_size,
shuffle=True,
num_workers=config.dataloader_num_workers
num_workers=config.dataloader_num_workers,
)
config.num_train_epochs = int(math.ceil(config.max_train_steps / len(train_dataloader)))
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
if config.max_train_steps is None:
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
print(f"--- Num samples = {len(train_dataset)}")
@@ -213,7 +194,7 @@ def train(config: TrainingConfig):
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
print(f"--- Total optimization steps = {config.max_train_steps}\n")
global_step = 0
last_save_step = 0
@@ -227,18 +208,20 @@ def train(config: TrainingConfig):
# Data tracking inits:
start_time, images_done = time.time(), 0
prompt_embeds_norms = {'main':[], 'reg':[]}
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': [], 'token_attention_loss': []}
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': []}
grad_norms, token_stds = {'unet': []}, {}
for i in range(len(text_encoders)):
grad_norms[f'text_encoder_{i}'] = []
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
# default values for cold (starting) optimizer lr:
base_unet_lr = 2.0e-4 if (config.is_lora and config.disable_ti) else 5.0e-5
if not config.is_lora:
base_unet_lr = 1.0e-5
# default value of cold (pre-warmup) optimizer lr:
if config.sd_model_version == "sdxl":
# let textual_inversion do the work first!
base_lr = 0.5e-5
elif config.sd_model_version == "sd15":
# let lora training kick in soonish
base_lr = 1.0e-4
#######################################################################################################
"""
@@ -252,8 +235,7 @@ def train(config: TrainingConfig):
)
optimizers = optimizer_collection.optimizers
if config.debug:
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
for epoch in range(config.num_train_epochs):
if config.aspect_ratio_bucketing:
@@ -266,12 +248,15 @@ def train(config: TrainingConfig):
completion_f = finegrained_epoch / config.num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training
if config.ti_optimizer != "prodigy" and optimizers['textual_inversion'] is not None:
# Apply the exponential learning rate
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 1.7
# Apply freezing condition
if completion_f > config.freeze_ti_after_completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] = 0.0
if config.ti_optimizer != "prodigy": # Update ti_learning rate gradually:
if optimizers['textual_inversion'] is not None:
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0
# warmup the ti-lr:
if config.ti_lr_warmup_steps > 0:
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0)
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
if config.freeze_ti_after_completion_f <= completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
if optimizers['text_encoders'] is not None:
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
@@ -283,26 +268,16 @@ def train(config: TrainingConfig):
if optimizers['unet'] is not None:
# Calculate the exponential factor
exp_factor = (config.unet_lr / base_unet_lr) ** (global_step / config.unet_lr_warmup_steps)
exp_factor = (config.unet_lr / base_lr) ** (global_step / config.unet_lr_warmup_steps)
# Apply the exponential learning rate
optimizers['unet'].param_groups[0]['lr'] = base_unet_lr * exp_factor
if completion_f < config.freeze_unet_before_completion_f:
optimizers['unet'].param_groups[0]['lr'] = 0.0
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
if not config.aspect_ratio_bucketing:
captions, vae_latent, mask = batch
else:
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
mask = mask.to(config.device)
captions = list(captions)
if config.caption_dropout > 0.0:
for i in range(len(captions)):
if np.random.rand() < config.caption_dropout:
captions[i] = config.token_dict["TOK"]
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
config, pipe, captions
)
@@ -334,23 +309,18 @@ def train(config: TrainingConfig):
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
return_dict=False,
)[0]
# Compute the loss:
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
losses['img_loss'].append(loss.item())
if not config.disable_ti:
token_attention_loss = compute_token_attention_loss(pipe, embedding_handler, captions, mask, daam_loss)
losses['token_attention_loss'].append(token_attention_loss.item())
loss = loss + config.token_attention_loss_w * token_attention_loss
if config.training_attributes["gpt_description"] and config.debug:
concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe)
# Dont apply this loss, just plot it for now:
loss += 0.0 * concept_description_loss
losses['concept_description_loss'].append(concept_description_loss.item())
if config.l1_penalty > 0.0 and unet_lora_parameters:
if config.l1_penalty > 0.0:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += config.l1_penalty * l1_norm
@@ -379,6 +349,12 @@ def train(config: TrainingConfig):
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
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()
#############################################################################################################
@@ -392,14 +368,9 @@ def train(config: TrainingConfig):
for std_i, std in enumerate(embedding_stds):
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
if global_step % 50 == 0 and not config.disable_ti and config.debug:
img_ratio = config.train_img_size[0] / config.train_img_size[1]
plot_token_attention_loss(config.output_dir, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio)
# Print some statistics:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)): #and global_step > 0:
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
if config.debug and (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > -1:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
config.save_as_json(
@@ -419,17 +390,27 @@ def train(config: TrainingConfig):
)
last_save_step = global_step
if config.debug:
embedding_handler.print_token_info()
if config.is_lora: # plotting this hist for full unet parameters can run OOM
plot_torch_hist(unet_lora_parameters, global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
#plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings()
for idx, text_encoder in enumerate(text_encoders):
if text_encoder is None:
continue
n = len(token_embeddings[f'txt_encoder_{idx}'])
for i in range(n):
token = trainable_tokens[f'txt_encoder_{idx}'][i]
# Strip any backslashes from the token name:
token = token.replace("/", "_")
embedding = token_embeddings[f'txt_encoder_{idx}'][i]
plot_torch_hist(embedding, global_step, os.path.join(config.output_dir, 'ti_embeddings') , f"enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'red')
embedding_handler.print_token_info()
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(
pipe = pipe,
render_size = config.validation_img_size,
@@ -439,8 +420,6 @@ def train(config: TrainingConfig):
is_lora = config.is_lora,
pretrained_model = config.pretrained_model,
lora_scale = config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs,
device = config.device,
checkpoint_folder = None
@@ -454,13 +433,14 @@ def train(config: TrainingConfig):
images_done += config.train_batch_size
global_step += 1
if global_step % (config.max_train_steps//100) == 0:
if global_step % (config.max_train_steps//20) == 0:
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))
if global_step > config.max_train_steps:
print("Reached max steps, stopping training!", flush = True)
print("Reached max steps, stopping training!")
break
# final_save
@@ -491,7 +471,8 @@ def train(config: TrainingConfig):
pretrained_model_version=config.pretrained_model["version"]
)
if config.debug and 0:
print("Running final inference round...")
if config.debug:
# Reload the entire pipe from disk + LoRa:
pipe_to_use = None
checkpoint_folder = output_save_dir
@@ -521,8 +502,6 @@ def train(config: TrainingConfig):
is_lora=config.is_lora,
pretrained_model=config.pretrained_model,
lora_scale=config.sample_imgs_lora_scale,
disable_ti = config.disable_ti,
prompt_modifier = config.prompt_modifier,
n_imgs = config.n_sample_imgs,
n_steps = 30,
device = config.device,
@@ -532,6 +511,13 @@ def train(config: TrainingConfig):
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"))
# 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:
print(f"Skipping final save, {output_save_dir} already exists")
@@ -545,8 +531,6 @@ def train(config: TrainingConfig):
config.job_time = time.time() - config.start_time
config.training_attributes["validation_prompts"] = validation_prompts
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
print("Training job complete, saving outputs...", flush = True)
print("------------------------------------------")
return config, output_save_dir
@@ -557,12 +541,7 @@ if __name__ == "__main__":
args = parser.parse_args()
config = TrainingConfig.from_json(file_path=args.config_filename)
print("Starting new LoRa training run with config:")
print(config)
print("------------------------------------------")
for progress in train(config=config):
print(f"Progress: {(100*progress):.2f}%", end="\r")
print("Training done :)")
print("Training done :)")
+2027
View File
File diff suppressed because it is too large Load Diff
+62 -95
View File
@@ -1,17 +1,22 @@
import os
import shutil
import tarfile
import json
import time
import random
import torch
import numpy as np
from PIL import Image
import pandas as pd
from dotenv import load_dotenv
from main import train
from trainer.config import TrainingConfig, model_paths
from trainer.utils.io import clean_filename
import folder_paths
import comfy.utils
from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig
from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything
class Eden_LoRa_trainer:
@classmethod
@@ -19,130 +24,92 @@ class Eden_LoRa_trainer:
return {
"required": {
"training_images_folder_path": ("STRING", {"default": "."}),
"mode": (["style", "face", "object"], {"default": "style"}),
"lora_name": ("STRING", {"default": "Eden_Token_LoRa"}),
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"training_resolution": ("INT", {"default": 512, "min": 256, "max": 1024}),
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 300, "min": 10, "max": 10000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0, "max": 0.005, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.0005, "min": 0.0, "max": 0.005, "step": 0.0001}),
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
"disable_ti": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 3, "min": 1, "max": 5}),
"save_checkpoint_every_n_steps": ("INT", {"default": 200, "min": 10, "max": 10000}),
"n_sample_imgs": ("INT", {"default": 4, "min": 2, "max": 10}),
"sample_imgs_lora_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.25}),
"plot_training_graphs_on_disk": ("BOOLEAN", {"default": False}),
"lora_name": ("STRING", {"default": ""}),
"sd_model_version": (["sdxl", "sd15"], ),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
"use_dora": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}),
}
}
CATEGORY = "Eden 🌱"
RETURN_TYPES = ("IMAGE", "STRING", "STRING", "STRING")
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
RETURN_TYPES = ("STRING",)
FUNCTION = "train_lora"
def train_lora(self,
training_images_folder_path,
ckpt_name,
lora_name,
mode,
training_resolution,
train_batch_size,
max_train_steps ,
ti_lr,
unet_lr,
lora_rank,
disable_ti,
n_tokens,
plot_training_graphs_on_disk,
save_checkpoint_every_n_steps,
n_sample_imgs,
sample_imgs_lora_scale,
seed,
def train_lora(self, training_images_folder_path,
name = lora_name,
concept_mode = "style",
sd_model_version = "sdxl",
seed = 0,
resolution = 521,
train_batch_size = 4,
max_train_steps = 400,
ti_lr = 0.001,
unet_lr = 0.001,
lora_rank = 16,
use_dora = False,
n_tokens = 2
):
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(
name=lora_name,
output_dir="output",
name="test",
lora_training_urls=training_images_folder_path,
concept_mode=mode,
ckpt_path=ckpt_path,
concept_mode=concept_mode,
sd_model_version=sd_model_version,
seed=seed,
resolution=training_resolution,
resolution=resolution,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=save_checkpoint_every_n_steps,
n_sample_imgs=(n_sample_imgs//2) * 2,
sample_imgs_lora_scale=sample_imgs_lora_scale,
checkpointing_steps=10000,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
use_dora=False,
use_dora=use_dora,
caption_model="blip",
disable_ti=disable_ti,
n_tokens=n_tokens,
verbose=True,
debug=plot_training_graphs_on_disk,
debug=True,
)
pbar = comfy.utils.ProgressBar(100)
with torch.inference_mode(False):
train_generator = train(config=config)
while True:
try:
progress_f = next(train_generator)
pbar.update_absolute(progress_f * 100)
except StopIteration as e:
config, output_save_dir = e.value # Capture the return value
break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
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['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
print(f"LORA training node finished in {config.job_time:.1f} seconds")
print("---------- Made with love by Eden.art 🌱 ----------")
# safetensors paths:
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
print(f"LORA training finished in {config.job_time:.1f} seconds")
print(f"Returning {out_path}")
# 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)
return (out_path,)
+18 -35
View File
@@ -14,6 +14,7 @@ from main import train
from typing import Iterator, Optional
from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig
from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything
@@ -65,19 +66,19 @@ class Predictor(BasePredictor):
),
max_train_steps: int = Input(
description="Number of training steps. Increasing this usually leads to overfitting, only viable if you have > 100 training imgs. For faces you may want to reduce to eg 300",
default=300
),
checkpointing_steps: int = Input(
description="Save a checkpoint every n steps (The final checkpoint will always be saved)",
default=10000
default=400
),
resolution: int = Input(
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 768",
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
default=512
),
train_batch_size: int = Input(
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
default=4
),
unet_lr: float = Input(
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
default=0.0003
default=0.001
),
ti_lr: float = Input(
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
@@ -87,25 +88,13 @@ class Predictor(BasePredictor):
description="Rank of LoRA embeddings for the unet.",
default=16
),
use_dora: bool = Input(
description="Use Dora instead of LoRa",
default=False,
),
n_tokens: int = Input(
description="How many new tokens to train (highly recommended to leave this at 2)",
ge=1, le=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
ge=1, le=3, default=2
),
seed: int = Input(
description="Random seed for reproducible training. Leave empty to use a random seed",
@@ -135,15 +124,13 @@ class Predictor(BasePredictor):
sd_model_version=sd_model_version,
seed=seed,
resolution=resolution,
validation_img_size=[validation_img_size, validation_img_size],
sample_imgs_lora_scale=sample_imgs_lora_scale,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=checkpointing_steps,
n_sample_imgs=n_sample_imgs,
checkpointing_steps=10000,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
use_dora=use_dora,
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
@@ -175,13 +162,9 @@ class Predictor(BasePredictor):
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
comfy_workflows_path = "ComfyUI_workflows"
if os.path.exists(comfy_workflows_path) and os.path.isdir(comfy_workflows_path):
for root, dirs, files in os.walk(comfy_workflows_path):
for file in files:
file_path = os.path.join(root, file)
arcname = os.path.relpath(file_path, os.path.dirname(comfy_workflows_path))
tar.add(file_path, arcname=arcname)
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
-16
View File
@@ -1,16 +0,0 @@
[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 = ""
+14 -23
View File
@@ -1,26 +1,17 @@
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
transformers>=4.38.1
diffusers>=0.27.2
ujson>=5.9.0
scipy>=1.12.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
triton<3.2.0
numpy>=1.26.4
opencv-python>=4.1.0.25
mediapipe>=0.10.11
openai>=1.14.0
python-dotenv
prodigyopt
omegaconf
ujson
+39 -41
View File
@@ -1,8 +1,10 @@
"""
Faces:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
Objects:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
@@ -12,12 +14,7 @@ https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets
Styles:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
https://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
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
"""
@@ -38,12 +35,12 @@ def hamming_distance(dict1, dict2):
#######################################################################################
# Setup the base experiment config:
exp_name = "ygor_sd15"
exp_name = "grimes"
caption_prefix = ""
mask_target_prompts = ""
n_exp = 100 # how many random experiment settings to generate
min_hamming_distance = 2 # min_n_params that have to be different from any previous experiment to be scheduled
nohup = False
n_exp = 200 # 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
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values
@@ -51,40 +48,43 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sd15"],
"sd_model_version": ["sd15", "sdxl"],
"lora_training_urls": [
"/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean"
"/home/rednax/Documents/datasets/grimes"
],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.9],
"caption_dropout": [0.2],
"concept_mode": ['face'],
"seed": [0],
"resolution": [512,640,768],
"train_batch_size": [8],
"n_sample_imgs": [8],
"max_train_steps": [2000],
"checkpointing_steps": [500],
"resolution": [512],
"train_batch_size": [4],
"n_sample_imgs": [6],
"max_train_steps": [400,800],
"checkpointing_steps": [100],
"gradient_accumulation_steps": [1],
"n_tokens": [3],
"disable_ti": ['true'],
"ti_lr": [0.0001],
"token_warmup_steps": [0],
"unet_lr": [0.001, 0.0003],
"lora_rank": [8,24,64],
"use_dora": ['false', 'true'],
"n_tokens": [2],
"ti_lr": [0.001,0.0005],
"ti_weight_decay": [0.001,0.0],
"l1_penalty": [0.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_optimizer_type": ['adamw'],
"is_lora": ['true'],
"unet_prodigy_growth_factor": [1.05],
"unet_lr": [0.001],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16,32],
"use_dora": ['false', 'true'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["florence", "blip", "no_caption"],
"augment_imgs_up_to_n": [40],
"caption_model": ["blip", "gpt4-v"],
"augment_imgs_up_to_n": [20,40],
"verbose": ['true'],
"debug": ['true']
}
@@ -100,7 +100,7 @@ shutil.rmtree(config_output_dir, ignore_errors=True)
os.makedirs(config_output_dir, exist_ok=True)
# Open the shell script file
try_sampling_n_times = 120
try_sampling_n_times = 200
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
resamples, combination = 0, None
@@ -122,6 +122,9 @@ for exp_index in tqdm(range(n_exp)): # number of combinations you want to gener
dirname = os.path.dirname(config_filename)
os.makedirs(dirname, exist_ok=True)
# Make some final adjustments to the experiment settings before saving to disk:
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
with open(config_filename, "w") as f:
json.dump(experiment_settings, f, indent=4)
break
@@ -143,12 +146,7 @@ def generate_sh_script(folder_path, output_sh_path):
# Write a command for each JSON file
for json_file in json_files:
file_path = os.path.join("scripts/", folder_path, json_file)
command = f"python main.py {file_path}\n"
if nohup:
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
command = f"python main.py {os.path.join(folder_path, json_file)}\n"
sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path)
-199
View File
@@ -1,199 +0,0 @@
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)
+59
View File
@@ -0,0 +1,59 @@
"""
Pre-trained checkpoint:
https://huggingface.co/stabilityai/stable-diffusion-3-medium-diffusers
"""
import os
import torch
from diffusers import StableDiffusion3Pipeline
import torch
# Load the pretrained model
pipe = StableDiffusion3Pipeline.from_pretrained(
"stabilityai/stable-diffusion-3-medium-diffusers",
torch_dtype=torch.float16,
seed = 0
)
# Load the LoRA weights from file
lora_weights_path = "sd3-xander/checkpoint-1000/pytorch_lora_weights.safetensors"
# Move model to GPU
pipe = pipe.to("cuda")
prompts = [
"This is a picture of a man holding a glass of beer. He is wearing a casual plaid shirt and jeans. The man is holding a frosty glass of golden beer with a thick, foamy head in his right hand, lifting it slightly as if making a toast. The background features wooden tables and chairs, vintage beer signs, and warm ambient lighting",
"A close up shot of a man as a dragon rider with a red sword named Za'roc. His face is clearly visible in the high cinematic shot.",
"A man in 2075, looking for the last drop of water in mars. 4k HDR",
# "A king in Skyrim"
]
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}_baseline.jpg"
)
)
pipe.load_lora_weights(lora_weights_path, alpha = 8)
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}.jpg"
)
)
print(f"Done!")
@@ -1,25 +0,0 @@
{
"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
@@ -1,20 +0,0 @@
{
"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
@@ -1,21 +0,0 @@
{
"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
}
@@ -1,21 +0,0 @@
{
"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
}
@@ -1,16 +0,0 @@
{
"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
@@ -1,21 +0,0 @@
{
"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
}
@@ -1,21 +0,0 @@
{
"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
}
@@ -1,21 +0,0 @@
{
"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
}
@@ -1,21 +0,0 @@
{
"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
@@ -1,16 +0,0 @@
{
"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
@@ -1,16 +0,0 @@
{
"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
}
+6 -37
View File
@@ -54,31 +54,10 @@ def set_adapter_scales(pipe, lora_scale = 1.0):
return pipe
import re
def remove_delimiter_characters(name: str, max_length: int = 255) -> str:
# Define a regular expression pattern to match all weird and special characters
pattern = r'[^\w.-]+' # Matches any character that is not alphanumeric, underscore, dot, or hyphen
# Replace all occurrences of the pattern with a single underscore
cleaned_name = re.sub(pattern, '_', name)
# Replace multiple consecutive underscores with a single underscore
cleaned_name = re.sub(r'_+', '_', cleaned_name)
# Strip leading or trailing underscores and dots
cleaned_name = cleaned_name.strip('_.')
# Ensure the name doesn't start with a dot (to avoid hidden files on Unix)
cleaned_name = cleaned_name.lstrip('.')
# Truncate to max_length if necessary
cleaned_name = cleaned_name[:max_length]
# Raise an error if the name is empty or malformed after cleaning
if not cleaned_name:
raise ValueError("Malformed name")
return cleaned_name
def remove_delimiter_characters(name: str):
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
# Convert to WebUI format
def convert_pytorch_lora_safetensors_to_webui(
@@ -156,7 +135,7 @@ def save_checkpoint(
embedding_handler.save_embeddings(
os.path.join(
output_dir,
f"{name}_{pretrained_model_version}_embeddings.safetensors"
f"{name}_embeddings.safetensors"
)
)
@@ -166,7 +145,7 @@ def save_checkpoint(
output_dir, "special_params.json"
)
)
if is_lora:
assert len(unet_lora_parameters) > 0, f"Expected len(unet_lora_parameters) to be greater than zero if is_lora is True"
@@ -205,21 +184,11 @@ def save_checkpoint(
convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_lora.safetensors")
output_filename=os.path.join(output_dir, f"{name}.safetensors")
)
else:
# Save the entire, finetuned unet weights:
unet.save_pretrained(save_directory = output_dir)
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
return
def load_checkpoint(
pretrained_model_version: str,
pretrained_model_path: str,
+34 -67
View File
@@ -3,47 +3,15 @@ from datetime import datetime
from pydantic import BaseModel
import json, time, os
from typing import Literal
from trainer.models import pretrained_models
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):
lora_training_urls: str
concept_mode: Literal["face", "style", "object"]
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
prompt_modifier: str = None # optional prompt modifier
caption_model: Literal["gpt4-v", "blip", "florence", "no_caption"] = "florence"
caption_dropout: float = 0.1 # dropout rate for captions: occasionally use empty prompt
sd_model_version: Literal["sdxl", "sd15", None] = None
ckpt_path: str = None # optional hardcoded checkpoint path
caption_model: Literal["gpt4-v", "blip"] = "blip"
sd_model_version: Literal["sdxl", "sd15", "sd3"]
pretrained_model: dict = None
seed: Union[int, None] = None
resolution: int = 512
@@ -51,36 +19,36 @@ class TrainingConfig(BaseModel):
train_img_size: List[int] = None
train_aspect_ratio: float = None
train_batch_size: int = 4
max_train_steps: int = 300
num_train_epochs: int = None
num_train_epochs: int = 10000
max_train_steps: int = 360
checkpointing_steps: int = 10000
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
unet_optimizer_type: Literal["adamw", "prodigy", "adamw_8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 0.0003
unet_lr: float = 1.0e-3
prodigy_d_coef: float = 1.0
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
lora_weight_decay: float = 0.004
ti_lr: float = 0.001
lora_weight_decay: float = 0.002
# if ti_lr is None, then we completely skip textual inversion
ti_lr: Union[float, None] = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
freeze_ti_after_completion_f: float = 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
freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done
token_attention_loss_w: float = 3e-7
cond_reg_w: float = 0.0e-5
tok_cond_reg_w: float = 0.0e-5
tok_cov_reg_w: float = 0. # regularizes the token covariance matrix wrt pretrained, normal tokens
l1_penalty: float = 0.03 # Makes the unet lora matrix more sparse
tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
snr_gamma: float = 5.0
lora_alpha_multiplier: float = 1.0
lora_rank: int = 16
lora_rank: int = 12
use_dora: bool = False
left_right_flip_augmentation: bool = True
@@ -91,17 +59,22 @@ class TrainingConfig(BaseModel):
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
n_sample_imgs: int = 4
name: str = None
output_dir: str = "eden_lora_training_runs"
output_dir: str = "lora_models/unnamed"
debug: bool = False
allow_tf32: bool = True
disable_ti: bool = False
skip_gpt_cleanup: bool = False
remove_ti_token_from_prompts: bool = False
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 3
inserting_list_tokens: List[str] = ["<s0>","<s1>","<s2>"]
token_dict: dict = {"TOK": "<s0><s1><s2>"}
n_tokens: int = 2
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
token_dict: dict = {"TOK": "<s0><s1>"}
device: str = "cuda:0"
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
crops_coords_top_left_h: int = 0
crops_coords_top_left_w: int = 0
do_cache: bool = True
unet_learning_rate: float = 1.0
lr_num_cycles: int = 1
lr_power: float = 1.0
sample_imgs_lora_scale: float = 0.65 # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0
training_attributes: dict = {}
aspect_ratio_bucketing: bool = False
@@ -120,19 +93,16 @@ class TrainingConfig(BaseModel):
def __init__(self, **data):
super().__init__(**data)
self.pretrained_model = pretrained_models[self.sd_model_version]
if not self.ckpt_path:
self.pretrained_model = pretrained_models[self.sd_model_version]
else:
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
# add some metrics to the foldername:
lora_str = "dora" if self.use_dora else "lora"
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
if not self.name:
self.name = os.path.basename(self.lora_training_urls)[:40]
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
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"
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}"
os.makedirs(self.output_dir, exist_ok=True)
if self.seed is None:
@@ -141,9 +111,6 @@ class TrainingConfig(BaseModel):
if self.unet_lr_warmup_steps is None:
self.unet_lr_warmup_steps = self.max_train_steps
if self.checkpointing_steps < 1:
self.checkpointing_steps = self.max_train_steps
if self.concept_mode == "face":
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
+31 -34
View File
@@ -1,12 +1,12 @@
import os
import torch
import numpy as np
from tqdm import tqdm
import pandas as pd
import PIL
from PIL import Image
from torch.utils.data import Dataset
from typing import Tuple, Dict, List
from tqdm import tqdm
def prepare_image(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
@@ -33,6 +33,7 @@ class PreprocessedDataset(Dataset):
data_dir: str,
pipe,
vae_encoder,
do_cache: bool = False,
size: List[int] = [512, 512],
text_dropout: float = 0.0,
aspect_ratio_bucketing: bool = False,
@@ -42,14 +43,13 @@ class PreprocessedDataset(Dataset):
super().__init__()
self.data_dir = data_dir
self.csv_path = os.path.join(data_dir, "captions.csv")
self.data = pd.read_csv(self.csv_path, dtype={"caption": str})
self.data = pd.read_csv(self.csv_path)
self.captions = self.data["caption"]
self.captions = self.captions.str.lower()
for key, value in substitute_caption_map.items():
self.captions = self.captions.str.replace(key.lower(), value)
self.captions = self.captions.fillna("")
self.image_path = self.data["image_path"]
if "mask_path" not in self.data.columns:
@@ -63,31 +63,23 @@ class PreprocessedDataset(Dataset):
self.text_dropout = text_dropout
self.size = size
# If the training data is small we can keep everything in memory, otherwise offload to disk
self.do_cache = True if len(self.data) < 500 else False
if self.do_cache:
print("Encoding latents, masks and captions and storing in memory...\n")
if do_cache:
print("Caching latents, masks and captions...\n")
self.vae_latents = []
self.masks = []
self.do_cache = True
for idx in tqdm(range(len(self.data))):
vae_latent, mask, _ = self._process(idx)
if len(self.data) < 25:
print(self.captions[idx])
vae_latent, mask = self._process(idx)
self.vae_latents.append(vae_latent)
self.masks.append(mask.detach())
self.masks.append(mask)
else: # Store the latents and masks on disk
print("Encoding latents, masks and captions and storing on disk...\n")
self.vae_latents = None
self.masks = None
for idx in tqdm(range(len(self.data))):
vae_latent, mask, image_path = self._process(idx)
torch.save(vae_latent, os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
torch.save(mask, os.path.join(self.data_dir, f"{idx}_mask.pt"))
del self.vae_encoder
torch.cuda.empty_cache()
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
del self.vae_encoder
else:
self.do_cache = False
if aspect_ratio_bucketing:
print("Using aspect ratio bucketing.")
@@ -109,9 +101,12 @@ class PreprocessedDataset(Dataset):
def get_aspect_ratio_bucketed_batch(self):
assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()"
indices, resolution = self.bucket_manager.get_batch()
print(f"Got bucket batch: {indices}, resolution: {resolution}")
tok1, tok2, vae_latents, masks = [], [], [], []
for idx in indices:
if self.tokenizer_2 is None:
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
else:
@@ -147,24 +142,28 @@ class PreprocessedDataset(Dataset):
image = PIL.Image.open(image_path).convert("RGB")
if bucketing_resolution is None:
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to(
dtype=self.vae_encoder.dtype
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
else:
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
dtype=self.vae_encoder.dtype
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
vae_latent = self.vae_encoder.encode(image.to(self.vae_encoder.device)).latent_dist
vae_latent = self.vae_encoder.encode(image).latent_dist
dummy_vae_latent = vae_latent.sample()
if self.mask_path is None:
mask = torch.ones_like(dummy_vae_latent, dtype=self.vae_encoder.dtype)
mask = torch.ones_like(
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
else:
mask_path = self.mask_path[idx]
mask_path = os.path.join(self.data_dir, mask_path)
mask = PIL.Image.open(mask_path)
mask = prepare_mask(mask, self.size[0], self.size[1]).to(dtype=self.vae_encoder.dtype)
mask = prepare_mask(mask, self.size[0], self.size[1]).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
)
mask_dtype = mask.dtype
mask = mask.float()
@@ -176,7 +175,7 @@ class PreprocessedDataset(Dataset):
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
return vae_latent, mask.squeeze(), image_path
return vae_latent, mask.squeeze()
def __getitem__(
self, idx: int, bucketing_resolution:tuple = None
@@ -184,12 +183,10 @@ class PreprocessedDataset(Dataset):
if self.do_cache:
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
return self.captions[idx], vae_latent.squeeze().detach(), self.masks[idx].detach()
else: # Load from disk:
vae_latent = torch.load(os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
return self.captions[idx], vae_latent.squeeze(), self.masks[idx]
else: # This code pathway has not been tested in a long time and might be broken
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution)
vae_latent = vae_latent.sample() * self.vae_scaling_factor
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()
return caption, vae_latent.squeeze(), mask
+123 -24
View File
@@ -9,6 +9,7 @@ from typing import List, Optional, Dict
from safetensors.torch import save_file, safe_open
import matplotlib.pyplot as plt
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
from transformers import T5EncoderModel
class TokenEmbeddingsHandler:
def __init__(self, text_encoders, tokenizers):
@@ -31,7 +32,10 @@ class TokenEmbeddingsHandler:
continue
# Directly accessing and modifying the original weights tensor
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.requires_grad_(True)
else:
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
def get_trainable_embeddings(self):
@@ -49,10 +53,22 @@ class TokenEmbeddingsHandler:
continue
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
if isinstance(text_encoder, T5EncoderModel):
indices_tensor = torch.tensor(
indices,
dtype=torch.long,
device=text_encoder.encoder.embed_tokens.weight.device
)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.encoder.embed_tokens.weight[indices_tensor]
else:
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
embeddings[f'txt_encoder_{idx}'] = token_embeddings
# Get all corresponding tokens for these embeddings
@@ -192,9 +208,19 @@ class TokenEmbeddingsHandler:
self.non_train_ids = all_indices[inu]
# random initialization of new tokens
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
"""
handle both T5EncoderModel and other text encoders
T5EncoderModel is present in sd3
"""
if isinstance(text_encoder, T5EncoderModel):
std_token_embedding = (
text_encoder.encoder.embed_tokens.weight.data.std(dim=1).mean()
)
else:
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
if starting_toks is not None:
@@ -207,14 +233,28 @@ class TokenEmbeddingsHandler:
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
else:
std_multiplier = 1.0
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
init_embeddings = torch.randn(len(self.train_ids), text_encoder.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
else:
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
current_std = init_embeddings.std(dim=1).mean()
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[self.train_ids] = init_embeddings.clone()
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
if isinstance(text_encoder, T5EncoderModel):
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.encoder.embed_tokens.weight.data.clone()
else:
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
@@ -260,7 +300,11 @@ class TokenEmbeddingsHandler:
# original_size = (config.resolution, config.resolution)
original_size = (1024, 1024)
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:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
@@ -393,6 +437,7 @@ class TokenEmbeddingsHandler:
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
optimizer_ti.step()
self.fix_embedding_std(config.off_ratio_power)
optimizer_ti.zero_grad()
if config.debug:
@@ -409,14 +454,27 @@ class TokenEmbeddingsHandler:
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
if isinstance(text_encoder, T5EncoderModel):
assert text_encoder.encoder.embed_tokens.weight.data.shape[
0
] == len(self.tokenizers[idx]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
]
)
else:
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
tensors[txt_encoder_keys[idx]] = new_token_embeddings
save_file(tensors, file_path)
@@ -425,6 +483,41 @@ class TokenEmbeddingsHandler:
def device(self):
return self.text_encoders[0].device
def fix_embedding_std(self, off_ratio_power=0.1):
if off_ratio_power == 0.0:
return
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
# Get the standard deviation target and current embeddings.
target_std = self.embeddings_settings[f"std_token_embedding_{idx}"]
embeddings, _ = self.get_trainable_embeddings()
new_embeddings = embeddings[f'txt_encoder_{idx}']
assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!"
new_stds = new_embeddings.std(dim=1)
#off_ratios = target_std.float() / new_stds.float()
off_ratios = target_std / new_stds
# Check if off_ratios are within an acceptable range.
if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1):
# Convert the pytorch tensor into a list of python floats:
off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3)
print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...")
# Adjust embeddings using the computed ratios.
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
index_updates = ~index_no_updates
multiplier_values = off_ratios**off_ratio_power
multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings)
text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values
idx += 1
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
# Assuming new tokens are of the format <s_i>
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
@@ -434,9 +527,15 @@ class TokenEmbeddingsHandler:
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
assert self.train_ids is not None, "New tokens could not be converted to IDs."
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
if not os.path.exists(file_path):
+6 -10
View File
@@ -155,7 +155,11 @@ def get_conditioning_signals(config, pipe, captions):
# original_size = (config.resolution, config.resolution)
original_size = (1024, 1024)
target_size = (config.resolution, config.resolution)
crops_coords_top_left = (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:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
@@ -241,7 +245,7 @@ def encode_prompt_advanced(
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
This allows interpolating the strength of the trained token in the final image.
"""
if lora_path and token_scale != 0:
if lora_path:
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
else:
lora_prompt = prompt
@@ -295,8 +299,6 @@ def render_images(
is_lora,
pretrained_model,
lora_scale,
disable_ti=False,
prompt_modifier=None,
n_steps=25,
n_imgs=4,
device="cuda:0",
@@ -326,9 +328,6 @@ def render_images(
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
validation_prompts_raw[0] = "<concept>"
if prompt_modifier:
validation_prompts_raw = [prompt_modifier.format(prompt) for prompt in validation_prompts_raw]
if (
checkpoint_folder is not None
): # reload the entire pipeline from disk and load in the lora module
@@ -377,7 +376,6 @@ def render_images(
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
)
pipeline_args["prompt_embeds"] = c
@@ -417,7 +415,6 @@ def render_images_eval(
pretrained_model: dict,
trigger_text: str,
lora_scale=0.7,
disable_ti=False,
n_steps=25,
n_imgs=4,
device="cuda:0",
@@ -472,7 +469,6 @@ def render_images_eval(
lora_scale,
guidance_scale=8,
concept_mode=concept_mode,
token_scale = 0 if disable_ti else None
)
pipeline_args["prompt_embeds"] = c
+15 -78
View File
@@ -4,81 +4,7 @@ import matplotlib.pyplot as plt
import torch
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
from trainer.inference import get_conditioning_signals
import torch.nn.functional as F
def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_loss, verbose=0):
"""
Custom loss function to regularize the attention maps of the token embeddings.
"""
masks = masks[:, 0].float()
img_ratio = masks.shape[-1] / masks.shape[-2]
att_L2_losses = []
ti_heatmaps = []
ti_masks = []
att_reg_threshold = 0.0
# attention_maps.shape = [n_layers, batch_size, w, h, 77]
attention_maps = daam_loss.process_and_stack_attention_scores(img_ratio)
n_layers, batch_size, w, h, n_tokens = attention_maps.shape
# reshape masks to match attention maps:
# masks.shape = [batch_size, w2, h2]
masks = F.interpolate(masks.unsqueeze(1), size=(attention_maps.shape[-3], attention_maps.shape[-2])).squeeze(1)
masks = masks.unsqueeze(0).unsqueeze(-1)
masks = masks.repeat(n_layers, 1, 1, 1, n_tokens)
for batch_index, caption in enumerate(captions):
token_indices = pipe.tokenizer.encode(caption)
# Penalize the mean attention score of each token:
mean_att_per_token = attention_maps[:,batch_index, :, :, 1:len(token_indices)-1].mean(dim=[0,1,2])
att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean()
att_L2_losses.append(att_L2_loss)
try:
ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids]
except:
continue
batch_ti_heatmaps, batch_ti_masks = [], []
# Extract the attention heatmaps corresponding to the trainable token embeddings:
for text_token_index in ti_token_indices:
ti_heatmap = attention_maps[:,batch_index, :, :, text_token_index].mean(dim=0)
ti_mask = masks[:,batch_index, :, :, text_token_index].mean(dim=0)
batch_ti_heatmaps.append(ti_heatmap.float())
batch_ti_masks.append(ti_mask)
ti_heatmaps.append(torch.stack(batch_ti_heatmaps))
ti_masks.append(torch.stack(batch_ti_masks))
if len(ti_heatmaps) == 0:
return torch.tensor(0.0).to(masks.dtype)
ti_heatmaps = torch.stack(ti_heatmaps)
ti_masks = torch.stack(ti_masks)
#ti_heatmaps.shape = [batch_size, n_tokens, w, h]
token_means = ti_heatmaps.mean(dim=[2,3])
token_attention_scores = token_means.var(dim=1)
# Avoid large attention scores in general:
reg_loss_0 = 5.0 * torch.stack(att_L2_losses).mean()
# Avoid large attention scores for ti tokens, inside the masked region:
reg_loss_1 = 1.0 * (torch.relu(ti_heatmaps * ti_masks)**2).mean()
# Avoid large attention scores for ti tokens, outside of the masked region:
reg_loss_2 = 2.0 * (torch.relu(ti_heatmaps * (1 - ti_masks) + 10)**2).mean()
# Make the Ti tokens have similar avg attention scores (equal distribution of concept information over tokens):
reg_loss_3 = 1.0 * token_attention_scores.mean()
if verbose:
print(f"reg_loss_0: {reg_loss_0.item():.4f}")
print(f"reg_loss_1: {reg_loss_1.item():.4f}")
print(f"reg_loss_2: {reg_loss_2.item():.4f}")
print(f"reg_loss_3: {reg_loss_3.item():.4f}")
return (reg_loss_0 + reg_loss_1 + reg_loss_2 + reg_loss_3).to(masks.dtype)
from transformers import T5EncoderModel
def compute_snr(noise_scheduler, timesteps):
"""
@@ -179,7 +105,14 @@ class ConditioningRegularizer:
def __init__(self, config, embedding_handler):
self.config = config
self.embedding_handler = embedding_handler
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
self.target_norms = {
"sdxl": 34.5,
"sd15": 27.8,
"sd3": 34.5
}
print(f'\033[91m[trainer.loss.ConditioningRegularizer] WARNING: Using a magic number: 34.5 for the target norm of sd3. We do not know if this is the ideal value. This might cause bugs or even break training completely.\033[0m')
self.target_norm = self.target_norms[config.sd_model_version]
self.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
@@ -189,11 +122,15 @@ class ConditioningRegularizer:
if tokenizer is None:
idx += 1
continue
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
if isinstance(text_encoder, T5EncoderModel):
pretrained_token_embeddings = text_encoder.encoder.embed_tokens.weight.data
else:
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
idx += 1
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.01, pipe=None):
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.003, pipe=None):
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
+54 -32
View File
@@ -4,30 +4,50 @@ import subprocess
import torch
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"
SD3_MODEL_CACHE = "models/stable-diffusion-3-medium"
#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"},
"sd3": {"path": SD3_MODEL_CACHE, "url": None, "version": "sd3"}
}
############################################################################################################
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:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
tokenizer_two, text_encoder_two = None, None
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
print(f"Loading model weights from {pretrained_model['path']} with dtype: {weight_dtype}...")
try:
print("Loading as SDXL model...")
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sdxl"
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
except:
print("Loading as SD15 model...")
if pretrained_model['version'] == "sd15":
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sd15"
else:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
print(f"Loaded {sd_model_version} model!")
pipe = pipe.to(device, dtype=weight_dtype)
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
@@ -37,11 +57,24 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16):
text_encoder_one = pipe.text_encoder
vae.requires_grad_(False)
vae.to(device, dtype=weight_dtype)
if keep_vae_float32:
vae.to(device, dtype=torch.float32)
else:
vae.to(device, dtype=weight_dtype)
if weight_dtype != torch.float32:
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but might not be for training..")
unet.to(device, dtype=weight_dtype)
text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype)
tokenizer_two = text_encoder_two = None
if pretrained_model['version'] == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
return (
pipe,
tokenizer_one,
@@ -51,7 +84,7 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16):
text_encoder_two,
vae,
unet,
), sd_model_version
)
def download_weights(url, dest):
start = time.time()
@@ -75,27 +108,16 @@ def download_weights(url, dest):
print(f"Downloading {url} took {time.time() - start} seconds")
def print_trainable_parameters(model, model_name=''):
def print_trainable_parameters(model, model_name = ''):
trainable_params = 0
all_param = 0
for name, param in model.named_parameters():
all_param += param.numel()
if param.requires_grad and "token_embedding" not in name:
trainable_params += param.numel()
def format_param_count(count):
if count < 1000:
return f"{count}"
elif count < 1_000_000:
return f"{count/1000:.1f}K"
else:
return f"{count/1_000_000:.1f}M"
line_delimiter = "#" * 80
print(line_delimiter)
print(
f"Trainable {model_name} params: {format_param_count(trainable_params)} "
f"|| All params: {format_param_count(all_param)} "
f"|| trainable = {100 * trainable_params / all_param:.2f}%"
f"Trainable {model_name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
)
print(line_delimiter)
+8 -42
View File
@@ -3,6 +3,11 @@ import torch
import prodigyopt
from typing import Iterable
def count_trainable_params(model):
return sum([
x.numel() for x in model.parameters() if x.requires_grad
])
def get_unet_optimizer(
prodigy_d_coef: float,
prodigy_growth_factor: float,
@@ -13,12 +18,9 @@ def get_unet_optimizer(
):
## unet_trainable_params can be unet.parameters() or a list of lora params
# These learning rates will get overwritten in main.py:
if optimizer_name == "adamw":
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
elif optimizer_name == "AdamW8bit":
import bitsandbytes as bnb
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
elif optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_unet = prodigyopt.Prodigy(
@@ -38,39 +40,6 @@ def get_unet_optimizer(
print(f"Created {optimizer_name} optimizer for unet!")
return optimizer_unet
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
def is_belong_to_blocks(key, blocks):
try:
for g in blocks:
if g in key:
return True
return False
except Exception as e:
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
if use_blora:
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
try:
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
attns = [
attn_processor_name.rsplit(".", 1)[0]
for attn_processor_name, _ in unet.attn_processors.items()
if is_belong_to_blocks(attn_processor_name, blocks)
]
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
return target_modules
except Exception as e:
raise type(e)(
f"failed to get_target_modules, due to: {e}. "
f"Please check the modules specified in --lora_unet_blocks are correct"
)
def get_unet_lora_parameters(
lora_rank,
lora_alpha_multiplier: float,
@@ -79,15 +48,12 @@ def get_unet_lora_parameters(
unet,
pipe,
):
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
unet_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian",
target_modules=target_modules,
target_modules=["to_k", "to_q", "to_v", "to_out.0", "conv2"],
#target_modules=["conv1", "conv2", "norm1", "norm2", "proj_in"], # TODO grid-search params for sd15
use_dora=use_dora,
)
+102 -137
View File
@@ -1,3 +1,7 @@
# Have SwinIR upsample
# Have BLIP auto caption
# Have CLIPSeg auto mask concept
import gc
import fnmatch
import mimetypes
@@ -21,7 +25,6 @@ import numpy as np
import pandas as pd
import torch
from tqdm import tqdm
from transformers import (
BlipForConditionalGeneration,
Blip2ForConditionalGeneration,
@@ -35,13 +38,13 @@ from transformers import (
from trainer.utils.io import download_and_prep_training_data
from trainer.utils.utils import fix_prompt
from trainer.config import model_paths
import re
import openai
from openai import OpenAI
from dotenv import load_dotenv
load_dotenv()
try:
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
client = OpenAI(api_key=OPENAI_API_KEY)
@@ -51,9 +54,11 @@ except:
client = None
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...)
MIN_GPT_PROMPTS = 3
MAX_GPT_PROMPTS = 80
MAX_GPT_PROMPTS = 50
def _find_files(pattern, dir="."):
"""Return list of files matching pattern in a given directory, in absolute format.
@@ -134,7 +139,7 @@ def swin_ir_sr(
"""
model = Swin2SRForImageSuperResolution.from_pretrained(
model_id, cache_dir = model_paths.get_path("SR")
model_id, cache_dir=MODEL_PATH
).to(device)
processor = Swin2SRImageProcessor()
@@ -188,9 +193,9 @@ def clipseg_mask_generator(
model = None
if any(target_prompts):
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = CLIPSegForImageSegmentation.from_pretrained(
model_id, cache_dir = model_paths.get_path("CLIP")
model_id, cache_dir=MODEL_PATH
).to(device)
masks = []
@@ -231,6 +236,7 @@ def clipseg_mask_generator(
return masks
import textwrap
def cleanup_prompts_with_chatgpt(
prompts,
@@ -331,12 +337,12 @@ def extract_gpt_concept_description(gpt_completion, concept_mode):
return concept_name
def post_process_captions(captions, text, concept_mode, job_seed, skip_gpt_cleanup=False):
def post_process_captions(captions, text, concept_mode, job_seed):
text = text.strip()
gpt_cleanup_worked = False
gpt_concept_description = None
if (len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client) and not skip_gpt_cleanup:
if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client:
retry_count = 0
while retry_count < 5:
try:
@@ -402,14 +408,14 @@ def blip_caption_dataset(
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "blip2" in model_id:
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = Blip2ForConditionalGeneration.from_pretrained(
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
).to(device)
else:
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = BlipForConditionalGeneration.from_pretrained(
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
).to(device)
for i, image in enumerate(tqdm(images)):
@@ -418,7 +424,6 @@ def blip_caption_dataset(
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65)
captions[i] = processor.decode(out[0], skip_special_tokens=True)
model.to("cpu")
del model
gc.collect()
torch.cuda.empty_cache()
@@ -440,6 +445,51 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
os.remove(output_path)
return base64_image
def gpt4_v_get_description(config, images):
if config.concept_mode == "object":
description = "object"
prompt = "Give a concise visual descriptioni of the object/figure/thing that all the grid-images have in common with at most 10 words. Dont start with statements like 'The image features...', just describe what you see."
elif config.concept_mode == "face":
description = "face"
prompt = "All the grid images depict a single person. Visually describe this person with at most 10 words. Dont start with statements like 'The image features...', just describe what you see. (eg an asian woman with long black hair)"
elif config.concept_mode == "style":
description = ""
prompt = "All these images share a common aesthetic style. Describe this style with at most 7 words. Dont start with statements like 'The image features...', just describe what you see. (eg impressionism collage surrealism)"
if not OPENAI_API_KEY:
print(f"Skipping GPT-4 Vision description because OPENAI_API_KEY is not set.")
return description
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {OPENAI_API_KEY}"
}
# TODO sample a grid img:
# .... TODO
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
payload = {
"model": "gpt-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(
images, captions,
batch_size=4,
@@ -450,7 +500,6 @@ def gpt4_v_caption_dataset(
return captions
prompt = "Concisely describe this image without assumptions with at most 20 words. Dont start with statements like 'The image features...', just describe what you see."
#prompt = "Concisely describe the main subject(s) in the image without assumptions with at most 20 words. Ignore the background and only describe the people / animals / objects in the foreground. Dont start with statements like 'The image features...', just describe what you see."
headers = {
"Content-Type": "application/json",
@@ -461,7 +510,7 @@ def gpt4_v_caption_dataset(
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
payload = {
"model": "gpt-4o",
"model": "gpt-4-turbo",
"messages": [
{
"role": "user",
@@ -499,84 +548,17 @@ def gpt4_v_caption_dataset(
device = "cuda:0" if torch.cuda.is_available() else "cpu"
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
@torch.no_grad()
def florence_caption_dataset(images, captions):
#workaround for unnecessary flash_attn requirement
from unittest.mock import patch
from transformers.dynamic_module_utils import get_imports
from transformers import AutoProcessor, AutoModelForCausalLM
def fixed_get_imports(filename: str | os.PathLike) -> list[str]:
if not str(filename).endswith("modeling_florence2.py"):
return get_imports(filename)
imports = get_imports(filename)
try:
imports.remove("flash_attn")
except:
pass
return imports
with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement
model = AutoModelForCausalLM.from_pretrained("microsoft/Florence-2-large", attn_implementation="sdpa", device_map=device, torch_dtype=torch_dtype,trust_remote_code=True)
processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large", trust_remote_code=True, cache_dir = model_paths.get_path("FLORENCE"))
for i, image in enumerate(tqdm(images)):
if captions[i] is None:
#prompt = random.choice(["<CAPTION>", "<DETAILED_CAPTION>"])
prompt = "<CAPTION>"
prompt = "<DETAILED_CAPTION>"
prompt = "<MORE_DETAILED_CAPTION>"
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
generated_ids = model.generate(
input_ids=inputs["input_ids"],
pixel_values=inputs["pixel_values"],
max_new_tokens=1024,
num_beams=random.choice([2,3,4])
)
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
parsed_answer = processor.post_process_generation(generated_text, task=prompt, image_size=(image.width, image.height))
caption = parsed_answer[prompt]
captions[i] = caption.replace("The image shows a ", "A ")
model.to('cpu')
del model
del processor
gc.collect()
torch.cuda.empty_cache()
return captions
@torch.no_grad()
def caption_dataset(
images: List[Image.Image],
captions: List[str],
caption_model: Literal["blip", "gpt4-v", "florence"] = "blip"
caption_model: Literal[str] = "blip"
) -> List[str]:
# if all captions are already generated, we dont need to do anything:
if all(captions):
print(f"All captions loaded from disk, skipping captioning...")
return captions
if "blip" in caption_model:
captions = blip_caption_dataset(images, captions)
elif "gpt4-v" in caption_model:
captions = gpt4_v_caption_dataset(images, captions)
elif "florence" in caption_model:
captions = florence_caption_dataset(images, captions)
else:
print("WARNING: not using any captions!")
captions = [""] * len(images)
gc.collect()
torch.cuda.empty_cache()
return captions
@@ -647,44 +629,20 @@ def random_crop(image, scale=(0.85, 0.95)):
return image.crop((left, top, left + new_width, top + new_height))
def gaussian_blur(image, radius = 1.0):
return image.filter(ImageFilter.GaussianBlur(radius=radius))
def gaussian_blur(image):
return image.filter(ImageFilter.GaussianBlur(radius=1))
def augment_image(image):
image = hue_augmentation(image)
image = color_jitter(image)
image = random_crop(image)
if random.random() < 0.5:
image = gaussian_blur(image, radius = random.uniform(0.0, 1.0))
image = gaussian_blur(image)
return image
def round_to_nearest_multiple(x, multiple):
return int(float(multiple) * round(float(x) / float(multiple)))
'''
For Stable Diffusion 1.5, outputs are optimised around 512x512 pixels. Many common fine-tuned versions of SD1.5 are optimised around 768x768. The best resolutions for common aspect ratios are typically:
1:1 (square): 512x512, 768x768
3:2 (landscape): 768x512
2:3 (portrait): 512x768
4:3 (landscape): 768x576
3:4 (portrait): 576x768
16:9 (widescreen): 912x512
9:16 (tall): 512x912
For SDXL, outputs are optimised around 1024x1024 pixels. The best resolutions for common aspect ratios are typically:
stable-diffusion-xl-1024-v0-9 supports generating images at the following dimensions:
1024 x 1024
1152 x 896
896 x 1152
1216 x 832
832 x 1216
1344 x 768
768 x 1344
1536 x 640
640 x 1536
'''
def calculate_new_dimensions(target_size, target_aspect_ratio):
"""
Calculate the new width and height given a target size and aspect ratio.
@@ -703,6 +661,8 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
return [new_width, new_height]
def load_and_save_masks_and_captions(
config,
concept_mode: str,
@@ -743,10 +703,9 @@ def load_and_save_masks_and_captions(
n_length = len(files)
files = sorted(files)[:n_length]
images, captions, img_paths = [], [], []
images, captions = [], []
for file in files:
images.append(load_image_with_orientation(file))
img_paths.append(file)
caption_file = os.path.splitext(file)[0] + ".txt"
if os.path.exists(caption_file) and use_dataset_captions:
with open(caption_file, "r") as f:
@@ -787,39 +746,41 @@ def load_and_save_masks_and_captions(
upscale_margin = 0.75
images = swin_ir_sr(images, target_size=(int(config.train_img_size[0]*upscale_margin), int(config.train_img_size[0]*upscale_margin)))
if add_lr_flips and len(images) < MAX_GPT_PROMPTS:
if add_lr_flips and len(images) < 40:
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
captions = captions + captions
# It's nice if we can achieve the gpt pass, so pre-augment the images if there's very few:
# Ensure we have at least 'augment_imgs_up_to_n' images through augmentation
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < MIN_GPT_PROMPTS:
print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images])
aug_caps.extend(captions)
print(f"Generating {len(images)} captions using {caption_model} in {concept_mode} mode...")
captions = caption_dataset(images, captions, caption_model = caption_model)
# Save captions back to disk:
for i, img_path in enumerate(img_paths):
caption_path = os.path.splitext(img_path)[0] + ".txt"
with open(caption_path, "w") as f:
f.write(captions[i])
images.extend(aug_imgs)
captions.extend(aug_caps)
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
images = images[:MAX_GPT_PROMPTS-1]
captions = captions[:MAX_GPT_PROMPTS-1]
if len(images) > 50 and caption_model != "blip":
print(f"Captioning a lot of ({len(images)}) images --> falling back to using blip!")
caption_model = "blip"
print(f"Generating {len(images)} captions using mode: {concept_mode}...")
captions = caption_dataset(images, captions, caption_model = caption_model)
# Cleanup prompts using chatgpt:
captions = [fix_prompt(caption) for caption in captions]
trigger_text = ""
gpt_concept_description = None
if not config.disable_ti:
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed, 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]
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
aug_imgs, aug_caps = [],[]
# if we still have a small amount of imgs, do some basic augmentation:
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images])
@@ -828,18 +789,19 @@ def load_and_save_masks_and_captions(
images.extend(aug_imgs)
captions.extend(aug_caps)
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
mask_target_prompts = gpt_concept_description
if mask_target_prompts is None or config.concept_mode == "style":
print("Disabling CLIP-segmentation")
mask_target_prompts = ""
temp = 999
else:
temp = config.clipseg_temperature
print(f"Generating {len(images)} masks...")
# Make sure we have a bias for the background pixels to never 100% ignore them
background_bias = 0.05
if not use_face_detection_instead:
@@ -892,11 +854,16 @@ def load_and_save_masks_and_captions(
os.remove(os.path.join(output_dir, file))
os.makedirs(output_dir, exist_ok=True)
if config.disable_ti:
# Make sure we've correctly inserted the TOK into every caption:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
for caption in captions:
print(caption)
if config.remove_ti_token_from_prompts:
print('------------------ WARNING -------------------')
print("Removing 'TOK, ' from captions...")
print("This will completely disable textual_inversion!!")
print("This will completely break textual_inversion!!")
print('------------------ WARNING -------------------')
if gpt_concept_description:
replace_str = gpt_concept_description
@@ -904,8 +871,6 @@ def load_and_save_masks_and_captions(
replace_str = ""
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
captions = [caption.replace("TOK", replace_str) for caption in captions]
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
# iterate through the images, masks, and captions and add a row to the dataframe for each
print("Saving final training dataset...")
-364
View File
@@ -1,364 +0,0 @@
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,18 +290,6 @@ def load_image_with_orientation(path, mode="RGB"):
elif orientation == 8:
image = image.rotate(90, expand=True)
if image.mode == 'P':
image = image.convert('RGBA')
if image.mode == 'CMYK':
image = image.convert('RGB')
# Remove alpha channel if present
if image.mode in ('RGBA', 'LA'):
background = Image.new('RGB', image.size, (255, 255, 255))
background.paste(image, mask=image.split()[3]) # 3 is the alpha channel
image = background
# Convert to the desired mode
return image.convert(mode)
+4 -13
View File
@@ -100,17 +100,15 @@ def print_system_info():
# Print disk space information
disk_usage = psutil.disk_usage('/')
total_disk = disk_usage.total // (1024 * 1024)
used_disk = disk_usage.used // (1024 * 1024)
free_disk = disk_usage.free // (1024 * 1024)
percent_disk_used = disk_usage.percent
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
print(f"Free disk space: {free_disk} MB with {percent_disk_used}% used")
# Print RAM information
virtual_mem = psutil.virtual_memory()
total_ram = virtual_mem.total // (1024 * 1024)
current_ram = virtual_mem.used // (1024 * 1024)
percent_ram_used = virtual_mem.percent
print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
print(f"Current used RAM: {current_ram} MB with {percent_ram_used}% used")
except Exception as e:
print(f'Error in gathering system info: {str(e)}')
@@ -124,13 +122,6 @@ def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1
# Flatten and concatenate all parameters into a single tensor
all_params = torch.cat([p.data.view(-1) for p in parameters])
# count number of parameters:
n_params = len(all_params)
if n_params == 0 or n_params > 1e9:
return
norm = torch.norm(all_params)
# Convert to CPU for plotting
@@ -237,7 +228,7 @@ def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dic
from scipy.signal import savgol_filter
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'):
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red', 'token_attention_loss': 'purple'}
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red'}
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
plot_smoothed = ['img_loss']
+2 -2
View File
@@ -2,9 +2,9 @@
val_prompts = {}
val_prompts['style'] = [
'a beautiful mountainous landscape, boulders, fresh water stream, setting sun',
'the stunning skyline of New York City, setting sun, skyscrapers, wallpaper',
'the stunning skyline of New York City',
'fruit hanging from a tree, highly detailed texture, soil, rain, drops, photo realistic, surrealism, highly detailed, 8k macrophotography',
'the Taj Mahal, stunning wallpaper, architecture, ancient, marble, white, intricate, detailed',
'the Taj Mahal, stunning wallpaper',
'A majestic tree rooted in circuits, leaves shimmering with data streams, stands as a beacon where the digital dawn caresses the fog-laden, binary soil—a symphony of pixels and chlorophyll.',
'A beautiful octopus, with swirling tendrils and a pulsating heart of fiery opal hues, hovers ethereally against a starry void, sculpted through a meticulous flame-working technique.',
'a stunning image of an aston martin sportscar',
+31
View File
@@ -0,0 +1,31 @@
{
"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": 7,
"n_sample_imgs": 6,
"max_train_steps": 5000,
"token_warmup_steps": 0,
"checkpointing_steps": 50,
"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": "blip",
"debug": true
}
+25
View File
@@ -0,0 +1,25 @@
{
"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
@@ -0,0 +1,29 @@
{
"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
@@ -0,0 +1,29 @@
{
"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
}