Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ca057f3597 | ||
|
|
9820c92bf3 | ||
|
|
c55440b95f | ||
|
|
d5f61bfd1d | ||
|
|
686a36126b | ||
|
|
59117c905f | ||
|
|
e83a7971e5 | ||
|
|
8fab8fe112 | ||
|
|
acb7c7bdc8 | ||
|
|
f676f708cc | ||
|
|
4df66e932d | ||
|
|
c1916918a1 | ||
|
|
176971db00 | ||
|
|
1568f10e6e | ||
|
|
3e35201c87 | ||
|
|
f621670210 | ||
|
|
2359aca6a9 | ||
|
|
cf8c0cd5e2 | ||
|
|
413e4d898c | ||
|
|
3b666a0b63 | ||
|
|
77a10d5912 | ||
|
|
6c3cf66d1f | ||
|
|
98fc2d293b | ||
|
|
bde4dc1992 | ||
|
|
92e3fe8599 | ||
|
|
ae4f6a0ee1 | ||
|
|
622303c48d | ||
|
|
7c0e1cbc62 | ||
|
|
e3f97bde50 | ||
|
|
cb03981dbe | ||
|
|
6568a8f124 | ||
|
|
09af837549 | ||
|
|
2be2e9b7b1 | ||
|
|
d5998e8fe6 | ||
|
|
c396f9cd73 | ||
|
|
e642930cea | ||
|
|
5798cbb187 | ||
|
|
acb40ea1d3 | ||
|
|
d263a95eb1 | ||
|
|
270e80364a | ||
|
|
869896466a | ||
|
|
2b7aff7b12 | ||
|
|
f3daa5bdeb | ||
|
|
414d43db2d | ||
|
|
bf067c8b98 | ||
|
|
9c154e8804 | ||
|
|
4b44d89722 | ||
|
|
c82df696a3 | ||
|
|
e04541805c | ||
|
|
9920a0ea5a | ||
|
|
14cb9ac335 | ||
|
|
4ca8695d86 | ||
|
|
22ac54a66d | ||
|
|
88946e88c2 | ||
|
|
3d5bf750ed | ||
|
|
e1bd34af91 | ||
|
|
87aa2e6872 | ||
|
|
682edd9333 | ||
|
|
2434d4846d | ||
|
|
c89b22d034 | ||
|
|
42e49afd78 | ||
|
|
5e4dcc2ece | ||
|
|
957420575e | ||
|
|
3f7c620f5a | ||
|
|
715afb1f65 | ||
|
|
6b6d6a6ad1 | ||
|
|
6eed725f94 | ||
|
|
6fa8dd47c1 | ||
|
|
2eda4dcba2 | ||
|
|
7627108bd6 | ||
|
|
9c63ec2d55 | ||
|
|
f2c1a42254 | ||
|
|
32291a3b2f | ||
|
|
590c8577d7 | ||
|
|
a83027d28c | ||
|
|
85facac79e | ||
|
|
ce9c608361 | ||
|
|
ff5cd50c49 | ||
|
|
e9555cd584 | ||
|
|
2e7879d39f | ||
|
|
c8d961dedf | ||
|
|
c0175fb67b | ||
|
|
cb98a8e5ab | ||
|
|
2de4cbdd17 | ||
|
|
345615c6be | ||
|
|
7693899e8b | ||
|
|
a4edb04deb | ||
|
|
d3d07bba0f | ||
|
|
2b7e174efd | ||
|
|
c542e73f5a | ||
|
|
c527fc690d | ||
|
|
c31b122686 | ||
|
|
e7aa728177 | ||
|
|
0a22b76511 | ||
|
|
6e044fb167 | ||
|
|
2ffb7ab1f5 | ||
|
|
8d553441d7 | ||
|
|
c9428c15d5 | ||
|
|
1eabc979a2 | ||
|
|
39d70c1bda | ||
|
|
4959c87ee7 | ||
|
|
cc467f04b7 | ||
|
|
f5b2569646 | ||
|
|
7f80ca88b3 | ||
|
|
5a2283b4c0 | ||
|
|
2c1bb4a3f9 | ||
|
|
ac2dcf787c | ||
|
|
ff165f32db | ||
|
|
8e71457f07 | ||
|
|
ef61d5f0e4 | ||
|
|
47a0d37b45 | ||
|
|
32836aa82f | ||
|
|
2b9f9c0f0a | ||
|
|
479465b47a | ||
|
|
f21efd7840 | ||
|
|
29855d94ba | ||
|
|
006a3b9750 | ||
|
|
b166f7c274 | ||
|
|
f6cfd6d0bc | ||
|
|
6a4fa1ef8d | ||
|
|
3cd17086f3 | ||
|
|
75137a7539 | ||
|
|
ee91f8c0c7 | ||
|
|
598c4204be | ||
|
|
8da6c91af8 | ||
|
|
01523350e1 | ||
|
|
c891203365 | ||
|
|
85599dc725 | ||
|
|
af16d1edcc | ||
|
|
b3db4ff3bb | ||
|
|
2a0761ae75 | ||
|
|
d8ee9ee522 | ||
|
|
4f084d251c | ||
|
|
fb521e9dbd | ||
|
|
f937106f6c | ||
|
|
5f8fe5c4c8 | ||
|
|
2f5eaeba7a | ||
|
|
7d8c845765 | ||
|
|
340336b53d | ||
|
|
5abed1487f | ||
|
|
b5cf857b1f | ||
|
|
e1813c75d1 | ||
|
|
4fde1b7dd6 | ||
|
|
892cf61c52 | ||
|
|
fdd531d3c5 | ||
|
|
ef6635fafb | ||
|
|
faf864ce15 | ||
|
|
b02221b23c | ||
|
|
d8b3175b57 |
+1
-8
@@ -22,11 +22,4 @@ datasets
|
|||||||
rendered_images
|
rendered_images
|
||||||
|
|
||||||
# trained models:
|
# trained models:
|
||||||
lora_models/*
|
lora_models/*
|
||||||
|
|
||||||
# Ignore the entire models folder by default:
|
|
||||||
models/*
|
|
||||||
|
|
||||||
### Include pipeline models: ###
|
|
||||||
!models/juggernaut_reborn.safetensors
|
|
||||||
!models/juggernaut_v6.safetensors
|
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
name: Publish to Comfy registry
|
||||||
|
on:
|
||||||
|
workflow_dispatch:
|
||||||
|
push:
|
||||||
|
branches:
|
||||||
|
- main
|
||||||
|
- master
|
||||||
|
paths:
|
||||||
|
- "pyproject.toml"
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
issues: write
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
publish-node:
|
||||||
|
name: Publish Custom Node to registry
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
if: ${{ github.repository_owner == 'edenartlab' }}
|
||||||
|
steps:
|
||||||
|
- name: Check out code
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
- name: Publish Custom Node
|
||||||
|
uses: Comfy-Org/publish-node-action@v1
|
||||||
|
with:
|
||||||
|
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||||
|
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||||
+5
-9
@@ -1,29 +1,25 @@
|
|||||||
data/
|
|
||||||
sd3_sweep_vis/
|
|
||||||
sd3_sweep_commands/
|
|
||||||
.ipynb_checkpoints/
|
|
||||||
cache
|
cache
|
||||||
__pycache__
|
__pycache__
|
||||||
|
.ipynb_checkpoints/
|
||||||
models
|
models
|
||||||
lora_models*
|
lora_models*
|
||||||
|
eden_lora_training_runs/
|
||||||
datasets
|
datasets
|
||||||
|
|
||||||
*.tar
|
*.tar
|
||||||
.env
|
.env
|
||||||
.cog
|
.cog
|
||||||
.huggingface
|
.huggingface
|
||||||
train.py
|
|
||||||
rendered_images*
|
rendered_images*
|
||||||
|
|
||||||
gridsearch*
|
gridsearch*
|
||||||
aesthetic_score_best_model.pth
|
aesthetic_score_best_model.pth
|
||||||
|
|
||||||
# experiment folders:
|
# experiment folders:
|
||||||
|
scripts/plots
|
||||||
conditioning_spaces/
|
conditioning_spaces/
|
||||||
training_args_x_*.json
|
training_args_x_*.json
|
||||||
xander_configs/
|
xander_configs/
|
||||||
debug/*
|
debug/*
|
||||||
wandb/
|
|
||||||
sd3_sweep_outputs/
|
|
||||||
sd3_face_sweep_configs/
|
|
||||||
|
|||||||
@@ -0,0 +1,782 @@
|
|||||||
|
{
|
||||||
|
"last_node_id": 26,
|
||||||
|
"last_link_id": 49,
|
||||||
|
"nodes": [
|
||||||
|
{
|
||||||
|
"id": 15,
|
||||||
|
"type": "Note Plus (mtb)",
|
||||||
|
"pos": {
|
||||||
|
"0": 301,
|
||||||
|
"1": -120,
|
||||||
|
"2": 0,
|
||||||
|
"3": 0,
|
||||||
|
"4": 0,
|
||||||
|
"5": 0,
|
||||||
|
"6": 0,
|
||||||
|
"7": 0,
|
||||||
|
"8": 0,
|
||||||
|
"9": 0
|
||||||
|
},
|
||||||
|
"size": [
|
||||||
|
519.1570279742359,
|
||||||
|
181.43570138536967
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 0,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [],
|
||||||
|
"outputs": [],
|
||||||
|
"title": "Unnamed",
|
||||||
|
"properties": {},
|
||||||
|
"widgets_values": [
|
||||||
|
"## How to trigger the LoRa:\nEden's trainer trains a token by default (if not disabled) which triggers the concept.\n\nFor \"object\" and \"face\" mode you just refer to your concept directly with the token, eg: \"a photo of embedding:MY\\_NAME\\_embedding\"\n\nFor \"style\" mode, you just prepend \"in the style of embedding:MY\\_NAME\\_embedding\" in the beginning of your prompt!",
|
||||||
|
"markdown",
|
||||||
|
"",
|
||||||
|
"one_dark"
|
||||||
|
],
|
||||||
|
"color": "#432",
|
||||||
|
"bgcolor": "#653",
|
||||||
|
"shape": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 13,
|
||||||
|
"type": "Note Plus (mtb)",
|
||||||
|
"pos": {
|
||||||
|
"0": -506,
|
||||||
|
"1": -142,
|
||||||
|
"2": 0,
|
||||||
|
"3": 0,
|
||||||
|
"4": 0,
|
||||||
|
"5": 0,
|
||||||
|
"6": 0,
|
||||||
|
"7": 0,
|
||||||
|
"8": 0,
|
||||||
|
"9": 0
|
||||||
|
},
|
||||||
|
"size": [
|
||||||
|
684.8613808966753,
|
||||||
|
209.8075855491436
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 1,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [],
|
||||||
|
"outputs": [],
|
||||||
|
"title": "Unnamed",
|
||||||
|
"properties": {},
|
||||||
|
"widgets_values": [
|
||||||
|
"## How to use:\n\n1. Find the training folder in ComfyUI/outputs\n2. Go into the checkpoints folder\n3. Pick the best checkpoint by looking at the validation_grid\n4. Copy the ***_embeddings.safetensors file to ComfyUI/models/embeddings\n5. Copy the ***_LoRa.safetensors file to ComfyUI/models/loras\n6. Hit Refresh in your ComfyUI\n7. Adjust this workflow to load both of those!\n8. Tweak the lora strength and embedding token strength to get the best results!\n\n\nHave fun! :)",
|
||||||
|
"markdown",
|
||||||
|
"",
|
||||||
|
"one_dark"
|
||||||
|
],
|
||||||
|
"color": "#432",
|
||||||
|
"bgcolor": "#653",
|
||||||
|
"shape": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 8,
|
||||||
|
"type": "VAEDecode",
|
||||||
|
"pos": [
|
||||||
|
1162,
|
||||||
|
188
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 140,
|
||||||
|
"1": 46
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 13,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "samples",
|
||||||
|
"type": "LATENT",
|
||||||
|
"link": 7
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "vae",
|
||||||
|
"type": "VAE",
|
||||||
|
"link": 8
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
9
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "VAEDecode"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 4,
|
||||||
|
"type": "CheckpointLoaderSimple",
|
||||||
|
"pos": [
|
||||||
|
-451,
|
||||||
|
162
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
429.578383119695,
|
||||||
|
98
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 2,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "MODEL",
|
||||||
|
"type": "MODEL",
|
||||||
|
"links": [
|
||||||
|
10
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "CLIP",
|
||||||
|
"type": "CLIP",
|
||||||
|
"links": [
|
||||||
|
12
|
||||||
|
],
|
||||||
|
"slot_index": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "VAE",
|
||||||
|
"type": "VAE",
|
||||||
|
"links": [
|
||||||
|
8
|
||||||
|
],
|
||||||
|
"slot_index": 2
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "CheckpointLoaderSimple"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"zavychromaxl_v90.safetensors"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 9,
|
||||||
|
"type": "SaveImage",
|
||||||
|
"pos": [
|
||||||
|
1327,
|
||||||
|
188
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 453.85968017578125,
|
||||||
|
"1": 407.6841125488281
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 14,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 9
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "SaveImage"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"ComfyUI"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 5,
|
||||||
|
"type": "EmptyLatentImage",
|
||||||
|
"pos": [
|
||||||
|
859,
|
||||||
|
26
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
271.419023112996,
|
||||||
|
106
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 3,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "LATENT",
|
||||||
|
"type": "LATENT",
|
||||||
|
"links": [
|
||||||
|
2
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "EmptyLatentImage"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
1024,
|
||||||
|
1024,
|
||||||
|
1
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 10,
|
||||||
|
"type": "LoraLoader",
|
||||||
|
"pos": [
|
||||||
|
31,
|
||||||
|
163
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 337.44464111328125,
|
||||||
|
"1": 126
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 6,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "model",
|
||||||
|
"type": "MODEL",
|
||||||
|
"link": 10
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "clip",
|
||||||
|
"type": "CLIP",
|
||||||
|
"link": 12
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "MODEL",
|
||||||
|
"type": "MODEL",
|
||||||
|
"links": [
|
||||||
|
24
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "CLIP",
|
||||||
|
"type": "CLIP",
|
||||||
|
"links": [
|
||||||
|
25,
|
||||||
|
26
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 1
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "LoraLoader"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"Eden_Token_LoRa_sdxl_LoRa.safetensors",
|
||||||
|
0.7000000000000001,
|
||||||
|
0.7000000000000001
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 3,
|
||||||
|
"type": "KSampler",
|
||||||
|
"pos": [
|
||||||
|
863,
|
||||||
|
186
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
267.5352926614355,
|
||||||
|
262
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 12,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "model",
|
||||||
|
"type": "MODEL",
|
||||||
|
"link": 24
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "positive",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"link": 40
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "negative",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"link": 44
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "latent_image",
|
||||||
|
"type": "LATENT",
|
||||||
|
"link": 2
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "LATENT",
|
||||||
|
"type": "LATENT",
|
||||||
|
"links": [
|
||||||
|
7
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "KSampler"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
1093427037772868,
|
||||||
|
"randomize",
|
||||||
|
30,
|
||||||
|
8,
|
||||||
|
"euler_ancestral",
|
||||||
|
"normal",
|
||||||
|
1
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 6,
|
||||||
|
"type": "CLIPTextEncode",
|
||||||
|
"pos": [
|
||||||
|
402,
|
||||||
|
124
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 422.84503173828125,
|
||||||
|
"1": 164.31304931640625
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 8,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "clip",
|
||||||
|
"type": "CLIP",
|
||||||
|
"link": 25
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "CONDITIONING",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"links": [
|
||||||
|
43
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "CLIPTextEncode"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"(in the style of embedding:Eden_Token_LoRa_sdxl_embeddings:1.0), \n\nwhat do i say to make me exist, oriental mythical beasts, in the golden danish age, in the history of television in the style of light violet and light red, serge najjar, playful and whimsical, associated press photo, afrofuturism-inspired, alasdair mclellan, electronic media"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 7,
|
||||||
|
"type": "CLIPTextEncode",
|
||||||
|
"pos": [
|
||||||
|
413,
|
||||||
|
389
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 425.27801513671875,
|
||||||
|
"1": 180.6060791015625
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 9,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "clip",
|
||||||
|
"type": "CLIP",
|
||||||
|
"link": 26
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "CONDITIONING",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"links": [
|
||||||
|
45
|
||||||
|
],
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "CLIPTextEncode"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
""
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 22,
|
||||||
|
"type": "ControlNetApplyAdvanced",
|
||||||
|
"pos": [
|
||||||
|
672,
|
||||||
|
741
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 166
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 11,
|
||||||
|
"mode": 4,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "positive",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"link": 43
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "negative",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"link": 45
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "control_net",
|
||||||
|
"type": "CONTROL_NET",
|
||||||
|
"link": 38
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 47
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "positive",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"links": [
|
||||||
|
40
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "negative",
|
||||||
|
"type": "CONDITIONING",
|
||||||
|
"links": [
|
||||||
|
44
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 1
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ControlNetApplyAdvanced"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
1,
|
||||||
|
0,
|
||||||
|
1
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 23,
|
||||||
|
"type": "ControlNetLoader",
|
||||||
|
"pos": [
|
||||||
|
255,
|
||||||
|
733
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 58
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 4,
|
||||||
|
"mode": 4,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "CONTROL_NET",
|
||||||
|
"type": "CONTROL_NET",
|
||||||
|
"links": [
|
||||||
|
38
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "ControlNetLoader"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"SDXL/controlnet-canny-sdxl-1.0/diffusion_pytorch_model_V2.safetensors"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 25,
|
||||||
|
"type": "AIO_Preprocessor",
|
||||||
|
"pos": [
|
||||||
|
258,
|
||||||
|
857
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 315,
|
||||||
|
"1": 82
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 7,
|
||||||
|
"mode": 4,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "image",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 48
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
47,
|
||||||
|
49
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "AIO_Preprocessor"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"CannyEdgePreprocessor",
|
||||||
|
512
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 26,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
679,
|
||||||
|
958
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
307.39388136076934,
|
||||||
|
34.317686638573264
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 10,
|
||||||
|
"mode": 4,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 49
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 24,
|
||||||
|
"type": "LoadImage",
|
||||||
|
"pos": [
|
||||||
|
-112,
|
||||||
|
857
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
315,
|
||||||
|
314
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 5,
|
||||||
|
"mode": 4,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "IMAGE",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
48
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "MASK",
|
||||||
|
"type": "MASK",
|
||||||
|
"links": null,
|
||||||
|
"shape": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "LoadImage"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"000050.jpg",
|
||||||
|
"image"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"links": [
|
||||||
|
[
|
||||||
|
2,
|
||||||
|
5,
|
||||||
|
0,
|
||||||
|
3,
|
||||||
|
3,
|
||||||
|
"LATENT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
7,
|
||||||
|
3,
|
||||||
|
0,
|
||||||
|
8,
|
||||||
|
0,
|
||||||
|
"LATENT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
8,
|
||||||
|
4,
|
||||||
|
2,
|
||||||
|
8,
|
||||||
|
1,
|
||||||
|
"VAE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
9,
|
||||||
|
8,
|
||||||
|
0,
|
||||||
|
9,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
10,
|
||||||
|
4,
|
||||||
|
0,
|
||||||
|
10,
|
||||||
|
0,
|
||||||
|
"MODEL"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
12,
|
||||||
|
4,
|
||||||
|
1,
|
||||||
|
10,
|
||||||
|
1,
|
||||||
|
"CLIP"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
24,
|
||||||
|
10,
|
||||||
|
0,
|
||||||
|
3,
|
||||||
|
0,
|
||||||
|
"MODEL"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
25,
|
||||||
|
10,
|
||||||
|
1,
|
||||||
|
6,
|
||||||
|
0,
|
||||||
|
"CLIP"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
26,
|
||||||
|
10,
|
||||||
|
1,
|
||||||
|
7,
|
||||||
|
0,
|
||||||
|
"CLIP"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
38,
|
||||||
|
23,
|
||||||
|
0,
|
||||||
|
22,
|
||||||
|
2,
|
||||||
|
"CONTROL_NET"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
40,
|
||||||
|
22,
|
||||||
|
0,
|
||||||
|
3,
|
||||||
|
1,
|
||||||
|
"CONDITIONING"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
43,
|
||||||
|
6,
|
||||||
|
0,
|
||||||
|
22,
|
||||||
|
0,
|
||||||
|
"CONDITIONING"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
44,
|
||||||
|
22,
|
||||||
|
1,
|
||||||
|
3,
|
||||||
|
2,
|
||||||
|
"CONDITIONING"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
45,
|
||||||
|
7,
|
||||||
|
0,
|
||||||
|
22,
|
||||||
|
1,
|
||||||
|
"CONDITIONING"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
47,
|
||||||
|
25,
|
||||||
|
0,
|
||||||
|
22,
|
||||||
|
3,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
48,
|
||||||
|
24,
|
||||||
|
0,
|
||||||
|
25,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
49,
|
||||||
|
25,
|
||||||
|
0,
|
||||||
|
26,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
]
|
||||||
|
],
|
||||||
|
"groups": [
|
||||||
|
{
|
||||||
|
"title": "Optional Controlnet",
|
||||||
|
"bounding": [
|
||||||
|
-180,
|
||||||
|
645,
|
||||||
|
1205,
|
||||||
|
536
|
||||||
|
],
|
||||||
|
"color": "#3f789e",
|
||||||
|
"font_size": 24,
|
||||||
|
"locked": false
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"config": {},
|
||||||
|
"extra": {
|
||||||
|
"ds": {
|
||||||
|
"scale": 0.683013455365071,
|
||||||
|
"offset": {
|
||||||
|
"0": 493.3539751601979,
|
||||||
|
"1": 181.8365096487169
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"version": 0.4
|
||||||
|
}
|
||||||
@@ -0,0 +1,312 @@
|
|||||||
|
{
|
||||||
|
"last_node_id": 21,
|
||||||
|
"last_link_id": 35,
|
||||||
|
"nodes": [
|
||||||
|
{
|
||||||
|
"id": 4,
|
||||||
|
"type": "Display Any (rgthree)",
|
||||||
|
"pos": [
|
||||||
|
899,
|
||||||
|
263
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 349.7635803222656,
|
||||||
|
"1": 88.69296264648438
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 5,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "source",
|
||||||
|
"type": "*",
|
||||||
|
"link": 34,
|
||||||
|
"dir": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "Display Any (rgthree)"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
""
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 3,
|
||||||
|
"type": "Display Any (rgthree)",
|
||||||
|
"pos": [
|
||||||
|
898,
|
||||||
|
150
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 347.36749267578125,
|
||||||
|
"1": 87.71485137939453
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 4,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "source",
|
||||||
|
"type": "*",
|
||||||
|
"link": 33,
|
||||||
|
"dir": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "Display Any (rgthree)"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
""
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 5,
|
||||||
|
"type": "Display Any (rgthree)",
|
||||||
|
"pos": [
|
||||||
|
897,
|
||||||
|
373
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 349.3501892089844,
|
||||||
|
"1": 98.81207275390625
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 6,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "source",
|
||||||
|
"type": "*",
|
||||||
|
"link": 35,
|
||||||
|
"dir": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "Display Any (rgthree)"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
""
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 2,
|
||||||
|
"type": "PreviewImage",
|
||||||
|
"pos": [
|
||||||
|
1307,
|
||||||
|
121
|
||||||
|
],
|
||||||
|
"size": {
|
||||||
|
"0": 557.6970825195312,
|
||||||
|
"1": 461.9648742675781
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 3,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [
|
||||||
|
{
|
||||||
|
"name": "images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"link": 32
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "PreviewImage"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 18,
|
||||||
|
"type": "Note Plus (mtb)",
|
||||||
|
"pos": {
|
||||||
|
"0": 411,
|
||||||
|
"1": -278,
|
||||||
|
"2": 0,
|
||||||
|
"3": 0,
|
||||||
|
"4": 0,
|
||||||
|
"5": 0,
|
||||||
|
"6": 0,
|
||||||
|
"7": 0,
|
||||||
|
"8": 0,
|
||||||
|
"9": 0
|
||||||
|
},
|
||||||
|
"size": {
|
||||||
|
"0": 1447.654541015625,
|
||||||
|
"1": 329.1551513671875
|
||||||
|
},
|
||||||
|
"flags": {},
|
||||||
|
"order": 0,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [],
|
||||||
|
"outputs": [],
|
||||||
|
"title": "Unnamed",
|
||||||
|
"properties": {},
|
||||||
|
"widgets_values": [
|
||||||
|
"\n## [Eden](https://www.eden.art/)\nThis LoRa trainer was made by the team behind https://www.eden.art/, led by https://x.com/xsteenbrugge\n\nIf you make awesome stuff w this trainer, give us a shout at:\nhttps://x.com/eden_art_ or\nhttps://www.instagram.com/eden.art____/\n\n---\n\n---\n\n## SD15 + SDXL:\nThis trainer works for both SD15 and SDXL models, but the default settings are primarily tuned for SDXL models. By selecting a specific ckpt_name, the trainer will automatically know if its an SDXL or SD15 model!\n\n---\n\n### A note on Embeddings:\nThis trainer optionally trains a textual inversion token into the LoRa, this is highly recommended when using SDXL models, but also means you have to load that token embedding when doing inference! See the example workflows in the repo.\n\n\nThe default settings work great for SDXL, SD15 usually need more training steps (eg 800) and sometimes benefits from disabling ti_training.\n\n---\n\n### A note on captioning:\nImages will get automatically captioned. It is recommended to put a .env file in the root of this custom node repo with your OpenAI API key, if found that will trigger a prompt_cleanup function that significantly improves results!\n",
|
||||||
|
"markdown",
|
||||||
|
"",
|
||||||
|
"one_dark"
|
||||||
|
],
|
||||||
|
"color": "#432",
|
||||||
|
"bgcolor": "#653",
|
||||||
|
"shape": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 19,
|
||||||
|
"type": "Note Plus (mtb)",
|
||||||
|
"pos": {
|
||||||
|
"0": 48,
|
||||||
|
"1": 106,
|
||||||
|
"2": 0,
|
||||||
|
"3": 0,
|
||||||
|
"4": 0,
|
||||||
|
"5": 0,
|
||||||
|
"6": 0,
|
||||||
|
"7": 0,
|
||||||
|
"8": 0,
|
||||||
|
"9": 0
|
||||||
|
},
|
||||||
|
"size": [
|
||||||
|
325.27622360123536,
|
||||||
|
609
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 1,
|
||||||
|
"mode": 0,
|
||||||
|
"inputs": [],
|
||||||
|
"outputs": [],
|
||||||
|
"title": "Unnamed",
|
||||||
|
"properties": {},
|
||||||
|
"widgets_values": [
|
||||||
|
"\n## A note on settings:\n\nIf you disbale\\_ti (not recommended) you'll get a normal LoRa that does not use a token embedding, in that case you typically need to train for a bit longer and increase the unet_lr.\n\n---\n\n\"training_images\" can both be a path to a local folder or a url to a public .zip file of imgs which will get downloaded.\n\n---\n\nYou can provide custom captions by placing a filename.txt file for each filename.jpg in the training_images folder\n\n---\n\nI highly recommend to keep the training resolution at either 512 or 768.\n\n---\n\nn_tokens = 1 is currently broken, need to fix that.\n\n---\n\nAn embedding + LoRa checkpoint will get saved every **save\\_checkpoint\\_every\\_n\\_steps**. Based on the sample image grid you can then pick the best checkpoint to use in your workflows!\n\n---\n\nSetting debug=True will save a bunch of additional graphs and visualizations to track whats happening during training for advanced users.",
|
||||||
|
"markdown",
|
||||||
|
"",
|
||||||
|
"one_dark"
|
||||||
|
],
|
||||||
|
"color": "#432",
|
||||||
|
"bgcolor": "#653",
|
||||||
|
"shape": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 21,
|
||||||
|
"type": "Eden_LoRa_trainer",
|
||||||
|
"pos": [
|
||||||
|
413,
|
||||||
|
124
|
||||||
|
],
|
||||||
|
"size": [
|
||||||
|
412.5926450093166,
|
||||||
|
591.4499899627035
|
||||||
|
],
|
||||||
|
"flags": {},
|
||||||
|
"order": 2,
|
||||||
|
"mode": 0,
|
||||||
|
"outputs": [
|
||||||
|
{
|
||||||
|
"name": "sample_images",
|
||||||
|
"type": "IMAGE",
|
||||||
|
"links": [
|
||||||
|
32
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 0
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "lora_path",
|
||||||
|
"type": "STRING",
|
||||||
|
"links": [
|
||||||
|
33
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 1
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "embedding_path",
|
||||||
|
"type": "STRING",
|
||||||
|
"links": [
|
||||||
|
34
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 2
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "final_msg",
|
||||||
|
"type": "STRING",
|
||||||
|
"links": [
|
||||||
|
35
|
||||||
|
],
|
||||||
|
"shape": 3,
|
||||||
|
"slot_index": 3
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"Node name for S&R": "Eden_LoRa_trainer"
|
||||||
|
},
|
||||||
|
"widgets_values": [
|
||||||
|
"https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
|
||||||
|
"style",
|
||||||
|
"Eden_Token_LoRa",
|
||||||
|
"zavychromaxl_v90.safetensors",
|
||||||
|
512,
|
||||||
|
4,
|
||||||
|
300,
|
||||||
|
0.001,
|
||||||
|
0.0005,
|
||||||
|
16,
|
||||||
|
false,
|
||||||
|
3,
|
||||||
|
200,
|
||||||
|
6,
|
||||||
|
0.7,
|
||||||
|
false,
|
||||||
|
40092,
|
||||||
|
"randomize"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"links": [
|
||||||
|
[
|
||||||
|
32,
|
||||||
|
21,
|
||||||
|
0,
|
||||||
|
2,
|
||||||
|
0,
|
||||||
|
"IMAGE"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
33,
|
||||||
|
21,
|
||||||
|
1,
|
||||||
|
3,
|
||||||
|
0,
|
||||||
|
"*"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
34,
|
||||||
|
21,
|
||||||
|
2,
|
||||||
|
4,
|
||||||
|
0,
|
||||||
|
"*"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
35,
|
||||||
|
21,
|
||||||
|
3,
|
||||||
|
5,
|
||||||
|
0,
|
||||||
|
"*"
|
||||||
|
]
|
||||||
|
],
|
||||||
|
"groups": [],
|
||||||
|
"config": {},
|
||||||
|
"extra": {
|
||||||
|
"ds": {
|
||||||
|
"scale": 0.6830134553650705,
|
||||||
|
"offset": [
|
||||||
|
106.80403405562511,
|
||||||
|
337.0121827057539
|
||||||
|
]
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"version": 0.4
|
||||||
|
}
|
||||||
+43
@@ -0,0 +1,43 @@
|
|||||||
|
Open Source Native License (OSNL)
|
||||||
|
Version 0.1 - March 1, 2024
|
||||||
|
|
||||||
|
Preamble
|
||||||
|
|
||||||
|
The Open Source Native License (OSNL) is designed to ensure that software remains free and open, fostering innovation and knowledge sharing within the community. It grants individuals, researchers, and commercial entities who open source their primary business assets, the freedom to use the software in any manner they choose.
|
||||||
|
|
||||||
|
This distinctive approach aims to balance the benefits of open-source development with the realities of commercial enterprise. It ensures that software remains a shared, community-driven resource while enabling businesses to thrive in an open-source ecosystem. Additional licenses are available for non-open source commercial entities.
|
||||||
|
|
||||||
|
1. Definitions
|
||||||
|
|
||||||
|
- "This License" refers to Version 1.0 of the Open Source Native License.
|
||||||
|
- "The Program" refers to the software distributed under this License.
|
||||||
|
- "You" refers to the individual or entity utilizing or contributing to the Program.
|
||||||
|
- "Primary Business Assets" are the core resources, capabilities, and technology that constitute the main value proposition and operational basis of your business.
|
||||||
|
|
||||||
|
2. Grant of License
|
||||||
|
|
||||||
|
Subject to the terms and conditions of this License, you are hereby granted a free, perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable license to use, reproduce, modify, distribute, and sublicense the Program, provided you comply with the following condition:
|
||||||
|
|
||||||
|
- Individual or researcher: You are granted the rights to use, modify, distribute, and contribute to the Program for any purpose, including educational, research, and personal projects, without the necessity to make your personal projects open source, provided these activities do not constitute a commercial enterprise. For any use that transitions to commercial purposes, the conditions applicable to commercial entities as outlined in this License will then apply.
|
||||||
|
- Commercial Entity who meets open source condition: Your primary business assets, including all core technologies, software, and platforms, must be available under an OSI-approved open source license or OSNL. This condition does not apply to ancillary or peripheral services not constituting primary business assets.
|
||||||
|
|
||||||
|
2.1 Commercial Use by Non-Open Source Businesses
|
||||||
|
|
||||||
|
Non-open source businesses that wish to utilize the Program or its derivatives as a component of their products or services are required to obtain an additional license. These entities must proactively contact Banodoco to request such a license. Banodoco reserves the right, at its own discretion, to grant or deny this additional license. Until an additional license is granted by Banodoco, non-open source businesses are not authorized to exercise any rights provided under this License regarding the use of the Program or its derivatives.
|
||||||
|
|
||||||
|
3. Redistribution
|
||||||
|
|
||||||
|
You may reproduce and distribute copies of the Program or derivative works thereof in any medium, with or without modifications, provided that you meet the following conditions:
|
||||||
|
|
||||||
|
- You must give any recipients of the Program a copy of this License.
|
||||||
|
- You must ensure that any modified files carry prominent notices stating that you changed the files.
|
||||||
|
- You must disclose the source of the Program, and if you distribute any portion of it in a compiled or object code form, you must also provide the full source code under this License.
|
||||||
|
- Any distribution of the Program or derivative works must comply with the Primary Business Open Source Condition.
|
||||||
|
|
||||||
|
4. Disclaimer of Warranty
|
||||||
|
|
||||||
|
THE PROGRAM IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES, OR OTHER LIABILITY ARISING FROM THE USE OF THE PROGRAM.
|
||||||
|
|
||||||
|
5. General
|
||||||
|
|
||||||
|
This License does not grant permission to use the trade names, trademarks, service marks, or product names of the Licensor, except as required for reasonable and customary use in describing the origin of the Program.
|
||||||
@@ -1,10 +1,11 @@
|
|||||||
# Trainer
|
# Trainer
|
||||||
|
|
||||||
This trainer was developed by the [**Eden** team](https://eden.art/)
|
This trainer was developed by the [**Eden** team](https://eden.art/), you can try our hosted version of the trainer in [**our app**](https://app.eden.art/).
|
||||||
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
|
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
|
||||||
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
|
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
|
||||||
|
|
||||||
The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
|
The outputs of this trainer are fully compatible with ComfyUI and AUTO111, see documentation [here](https://docs.eden.art/docs/guides/concepts/#exporting-loras-for-use-in-other-tools).
|
||||||
|
A full guide on training can be found in [**our docs**](https://docs.eden.art/docs/guides/concepts/#training).
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<strong>Training images:</strong><br>
|
<strong>Training images:</strong><br>
|
||||||
@@ -16,11 +17,26 @@ The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
|
|||||||
</p>
|
</p>
|
||||||
|
|
||||||
|
|
||||||
The trainer supports 3 default modes:
|
### The trainer can be run in 4 different ways:
|
||||||
|
- [**as a hosted service on our website**](https://app.eden.art/)
|
||||||
|
- [**as a hosted service through replicate**](https://replicate.com/edenartlab/sdxl-lora-trainer)
|
||||||
|
- **as a ComfyUI node**
|
||||||
|
- **as a standalone python script**
|
||||||
|
|
||||||
|
### Using in ComfyUI:
|
||||||
|
- Example workflows for how to run the trainer and do inference with it can be found in `/ComfyUI_workflows`
|
||||||
|
- Importantly this trainer uses a chatgpt call to cleanup the auto-generated prompts and inject the trainable token, this will only work if you have a .env file containing your OPENAI key in the root of the repo dir that contains a single line: `OPENAI_API_KEY=your_key_string` Everything will work without this, but results will be better if you set this up, especially for 'face' and 'object' modes.
|
||||||
|
|
||||||
|
### The trainer supports 3 default modes:
|
||||||
- **style**: used for learning the aesthetic style of a collection of images.
|
- **style**: used for learning the aesthetic style of a collection of images.
|
||||||
- **face**: used for learning a specific face (can be human, character, ...).
|
- **face**: used for learning a specific face (can be human, character, ...).
|
||||||
- **object**: will learn a specific object or thing featured in the training images.
|
- **object**: will learn a specific object or thing featured in the training images.
|
||||||
|
|
||||||
|
<p align="center">
|
||||||
|
<strong>Style training example:</strong><br>
|
||||||
|
<img src="assets/style_training_example.jpg" alt="Image 1" style="width:80%;"/>
|
||||||
|
</p>
|
||||||
|
|
||||||
## Setup
|
## Setup
|
||||||
|
|
||||||
Install all dependencies using
|
Install all dependencies using
|
||||||
@@ -29,7 +45,7 @@ Install all dependencies using
|
|||||||
|
|
||||||
then you can simply run:
|
then you can simply run:
|
||||||
|
|
||||||
`python main.py -c training_args.json`
|
`python main.py train_configs/training_args.json`
|
||||||
to start a training job.
|
to start a training job.
|
||||||
|
|
||||||
Adjust the arguments inside `training_args.json` to setup a custom training job.
|
Adjust the arguments inside `training_args.json` to setup a custom training job.
|
||||||
@@ -44,61 +60,25 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
|
|||||||
sudo chmod +x /usr/local/bin/cog
|
sudo chmod +x /usr/local/bin/cog
|
||||||
```
|
```
|
||||||
|
|
||||||
2. Build the image with `sudo cog build`
|
2. Build the image with `cog build`
|
||||||
3. Run a training run with `sudo sh cog_test_train.sh`
|
3. Run a training run with `sh cog_test_train.sh`
|
||||||
|
4. You can also go into the container with `cog run /bin/bash`
|
||||||
## Automatic Checkpoint Evaluation
|
|
||||||
|
|
||||||
This script uses CLIP img/txt similarity scores to evaluate how good the LoRa is vs how overfit.
|
|
||||||
Download the aesthetic predictor model checkpoint first from google drive. This should give you a file named: `aesthetic_score_best_model.pth` (99.2 MB)
|
|
||||||
|
|
||||||
```bash
|
|
||||||
gdown 1thEIlXVc8lkULVUBY9Ab45tsOERxkjxns
|
|
||||||
```
|
|
||||||
|
|
||||||
Once the model is downloaded, you can run the eval script with the following CLI args:
|
|
||||||
|
|
||||||
- `output_folder`: this is where the outputs of the model get saved as jpeg files
|
|
||||||
- `lora_path`: path to your LoRA checkpoint (make sure you edit `path_to_your_model_checkpoints` to point to the correct folder. It generally ends with something like `checkpoint-600` where `600` was the training step)
|
|
||||||
- `output_json`: save all scores in this json file
|
|
||||||
- `config_filename`: config file used for training
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python3 evaluate.py \
|
|
||||||
--output_folder eval_images \
|
|
||||||
--lora_path path_to_your_model_checkpoint \
|
|
||||||
--output_json eval_results.json \
|
|
||||||
--config_filename training_args.json
|
|
||||||
```
|
|
||||||
|
|
||||||
|
## Full unet finetuning
|
||||||
|
When running this trainer in native python, you can also perform full unet finetuning using something like (adjust to your needs)
|
||||||
|
`python main.py train_configs/full_finetuning_example.json`
|
||||||
|
|
||||||
## TODO's
|
## TODO's
|
||||||
|
|
||||||
Bugs:
|
Bugs:
|
||||||
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
|
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
|
||||||
|
- figure out why training is 3x slower through comfyui node versus just running main.py as a python job..?
|
||||||
Algo:
|
|
||||||
- Improve some of the chatgpt functionality:
|
|
||||||
- separate the "gpt_description" / "gpt_segmentation" prompt calls and make them run on a subset of prompts in case there's a lot of imgs / prompts (possibly use img_grids for some gpt4-v calls)
|
|
||||||
- currently some sub-optimal stuff can happen in preprocess() when there's less than 3 or more than 45 imgs, try to improve this
|
|
||||||
- Test if timesteps = torch.randint() can be improved: look at sdxl training code! (see https://github.com/huggingface/diffusers/blob/main/examples/advanced_diffusion_training/train_dreambooth_lora_sdxl_advanced.py#L1263, https://arxiv.org/pdf/2206.00364.pdf)
|
|
||||||
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
|
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
|
||||||
- test if textual inversion training can also happen with prodigy_optimizer
|
|
||||||
- improve data augmentation, eg by adding outpainted, smaller versions of faces / objects
|
|
||||||
|
|
||||||
Small, minor tweaks:
|
|
||||||
- preprocess.py: the imgs are first auto-captioned and then cropped, this is not ideal, swap this around!
|
|
||||||
|
|
||||||
Bigger improvements:
|
Bigger improvements:
|
||||||
- add stronger token regularization (eg CelebBasis spanning basis):
|
- integrate Flux / SD3
|
||||||
- Add multi-token training
|
- Add multi-concept training (multiple things represented by multiple tokens, trained into a single LoRa)
|
||||||
|
- add stronger token regularization (eg CelebBasis spanning basis)
|
||||||
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
|
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
|
||||||
- implement prompt-aligned: https://prompt-aligned.github.io/
|
- implement prompt-aligned: https://prompt-aligned.github.io/
|
||||||
|
|
||||||
Tuning Experiments:
|
|
||||||
- try-out conditioning noise injection during training to increase robustness
|
|
||||||
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
|
|
||||||
but it then recovers. Can we avoid this collapse? Is the learning rate too high?
|
|
||||||
- offset noise
|
|
||||||
- AB test Dora vs Lora
|
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 430 KiB |
@@ -3,17 +3,14 @@
|
|||||||
|
|
||||||
build:
|
build:
|
||||||
gpu: true
|
gpu: true
|
||||||
cuda: "11.8"
|
cuda: "12.1"
|
||||||
python_version: "3.9"
|
python_version: "3.11"
|
||||||
system_packages:
|
system_packages:
|
||||||
- "ffmpeg"
|
- "ffmpeg"
|
||||||
- "libgl1-mesa-glx"
|
|
||||||
- "libegl1-mesa-dev"
|
|
||||||
- "libsm6"
|
|
||||||
- "libxext6"
|
|
||||||
python_requirements: requirements.txt
|
python_requirements: requirements.txt
|
||||||
run:
|
run:
|
||||||
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task
|
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
|
||||||
|
|
||||||
predict: "predict.py:Predictor"
|
predict: "predict.py:Predictor"
|
||||||
image: "r8.im/edenartlab/sdxl-lora-trainer"
|
image: "r8.im/edenartlab/sdxl-lora-trainer"
|
||||||
|
|||||||
+7
-6
@@ -1,12 +1,13 @@
|
|||||||
# Set GPU ID to run these jobs on:
|
# Set GPU ID to run these jobs on:
|
||||||
GPU_ID="device=2"
|
GPU_ID="device=3"
|
||||||
|
|
||||||
cog predict --gpus $GPU_ID \
|
cog predict --gpus $GPU_ID \
|
||||||
-i name="xander_sdxl_cog" \
|
-i name="xander_sdxl_cog" \
|
||||||
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip" \
|
-i lora_training_urls="https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip" \
|
||||||
-i concept_mode="face" \
|
-i concept_mode="style" \
|
||||||
-i sd_model_version="sdxl" \
|
-i sd_model_version="sdxl" \
|
||||||
-i max_train_steps="360" \
|
-i max_train_steps="300" \
|
||||||
-i caption_model="blip" \
|
-i sample_imgs_lora_scale="0.7" \
|
||||||
-i debug="False" \
|
-i n_sample_imgs="6" \
|
||||||
|
-i debug="True" \
|
||||||
-i seed="0"
|
-i seed="0"
|
||||||
@@ -1,533 +0,0 @@
|
|||||||
{
|
|
||||||
"last_node_id": 12,
|
|
||||||
"last_link_id": 23,
|
|
||||||
"nodes": [
|
|
||||||
{
|
|
||||||
"id": 7,
|
|
||||||
"type": "CLIPTextEncode",
|
|
||||||
"pos": [
|
|
||||||
413,
|
|
||||||
389
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 425.27801513671875,
|
|
||||||
"1": 180.6060791015625
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 6,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "clip",
|
|
||||||
"type": "CLIP",
|
|
||||||
"link": 16
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "CONDITIONING",
|
|
||||||
"type": "CONDITIONING",
|
|
||||||
"links": [
|
|
||||||
6
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "CLIPTextEncode"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
"text, watermark"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 8,
|
|
||||||
"type": "VAEDecode",
|
|
||||||
"pos": [
|
|
||||||
1209,
|
|
||||||
188
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 210,
|
|
||||||
"1": 46
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 8,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "samples",
|
|
||||||
"type": "LATENT",
|
|
||||||
"link": 7
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "vae",
|
|
||||||
"type": "VAE",
|
|
||||||
"link": 8
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "IMAGE",
|
|
||||||
"type": "IMAGE",
|
|
||||||
"links": [
|
|
||||||
9
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "VAEDecode"
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 12,
|
|
||||||
"type": "Reroute",
|
|
||||||
"pos": [
|
|
||||||
220,
|
|
||||||
14
|
|
||||||
],
|
|
||||||
"size": [
|
|
||||||
75,
|
|
||||||
26
|
|
||||||
],
|
|
||||||
"flags": {},
|
|
||||||
"order": 3,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "",
|
|
||||||
"type": "*",
|
|
||||||
"link": 22
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "",
|
|
||||||
"type": "MODEL",
|
|
||||||
"links": [
|
|
||||||
19
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"showOutputText": false,
|
|
||||||
"horizontal": false
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 3,
|
|
||||||
"type": "KSampler",
|
|
||||||
"pos": [
|
|
||||||
863,
|
|
||||||
186
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 315,
|
|
||||||
"1": 262
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 7,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "model",
|
|
||||||
"type": "MODEL",
|
|
||||||
"link": 19
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "positive",
|
|
||||||
"type": "CONDITIONING",
|
|
||||||
"link": 4
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "negative",
|
|
||||||
"type": "CONDITIONING",
|
|
||||||
"link": 6
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "latent_image",
|
|
||||||
"type": "LATENT",
|
|
||||||
"link": 2
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "LATENT",
|
|
||||||
"type": "LATENT",
|
|
||||||
"links": [
|
|
||||||
7
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "KSampler"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
1,
|
|
||||||
"fixed",
|
|
||||||
25,
|
|
||||||
8,
|
|
||||||
"euler",
|
|
||||||
"normal",
|
|
||||||
1
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 5,
|
|
||||||
"type": "EmptyLatentImage",
|
|
||||||
"pos": [
|
|
||||||
473,
|
|
||||||
609
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 315,
|
|
||||||
"1": 106
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 0,
|
|
||||||
"mode": 0,
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "LATENT",
|
|
||||||
"type": "LATENT",
|
|
||||||
"links": [
|
|
||||||
2
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "EmptyLatentImage"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
768,
|
|
||||||
768,
|
|
||||||
1
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 4,
|
|
||||||
"type": "CheckpointLoaderSimple",
|
|
||||||
"pos": [
|
|
||||||
-467,
|
|
||||||
120
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 315,
|
|
||||||
"1": 98
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 1,
|
|
||||||
"mode": 0,
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "MODEL",
|
|
||||||
"type": "MODEL",
|
|
||||||
"links": [
|
|
||||||
10
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "CLIP",
|
|
||||||
"type": "CLIP",
|
|
||||||
"links": [
|
|
||||||
12
|
|
||||||
],
|
|
||||||
"slot_index": 1
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "VAE",
|
|
||||||
"type": "VAE",
|
|
||||||
"links": [
|
|
||||||
8
|
|
||||||
],
|
|
||||||
"slot_index": 2
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "CheckpointLoaderSimple"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
"juggernaut_reborn.safetensors"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 6,
|
|
||||||
"type": "CLIPTextEncode",
|
|
||||||
"pos": [
|
|
||||||
415,
|
|
||||||
186
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 422.84503173828125,
|
|
||||||
"1": 164.31304931640625
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 5,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "clip",
|
|
||||||
"type": "CLIP",
|
|
||||||
"link": 15
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "CONDITIONING",
|
|
||||||
"type": "CONDITIONING",
|
|
||||||
"links": [
|
|
||||||
4
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "CLIPTextEncode"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
"a photo of embedding:xander_sd15_embedding on the beach "
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 11,
|
|
||||||
"type": "Reroute",
|
|
||||||
"pos": [
|
|
||||||
220,
|
|
||||||
49
|
|
||||||
],
|
|
||||||
"size": [
|
|
||||||
75,
|
|
||||||
26
|
|
||||||
],
|
|
||||||
"flags": {},
|
|
||||||
"order": 4,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "",
|
|
||||||
"type": "*",
|
|
||||||
"link": 23
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "",
|
|
||||||
"type": "CLIP",
|
|
||||||
"links": [
|
|
||||||
15,
|
|
||||||
16
|
|
||||||
],
|
|
||||||
"slot_index": 0
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"showOutputText": false,
|
|
||||||
"horizontal": false
|
|
||||||
}
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 9,
|
|
||||||
"type": "SaveImage",
|
|
||||||
"pos": [
|
|
||||||
347,
|
|
||||||
-270
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 210,
|
|
||||||
"1": 270
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 9,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "images",
|
|
||||||
"type": "IMAGE",
|
|
||||||
"link": 9
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {},
|
|
||||||
"widgets_values": [
|
|
||||||
"ComfyUI"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"id": 10,
|
|
||||||
"type": "LoraLoader",
|
|
||||||
"pos": [
|
|
||||||
-72,
|
|
||||||
-229
|
|
||||||
],
|
|
||||||
"size": {
|
|
||||||
"0": 254.95774841308594,
|
|
||||||
"1": 127.86701202392578
|
|
||||||
},
|
|
||||||
"flags": {},
|
|
||||||
"order": 2,
|
|
||||||
"mode": 0,
|
|
||||||
"inputs": [
|
|
||||||
{
|
|
||||||
"name": "model",
|
|
||||||
"type": "MODEL",
|
|
||||||
"link": 10
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "clip",
|
|
||||||
"type": "CLIP",
|
|
||||||
"link": 12
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"outputs": [
|
|
||||||
{
|
|
||||||
"name": "MODEL",
|
|
||||||
"type": "MODEL",
|
|
||||||
"links": [
|
|
||||||
22
|
|
||||||
],
|
|
||||||
"shape": 3,
|
|
||||||
"slot_index": 0
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"name": "CLIP",
|
|
||||||
"type": "CLIP",
|
|
||||||
"links": [
|
|
||||||
23
|
|
||||||
],
|
|
||||||
"shape": 3,
|
|
||||||
"slot_index": 1
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"properties": {
|
|
||||||
"Node name for S&R": "LoraLoader"
|
|
||||||
},
|
|
||||||
"widgets_values": [
|
|
||||||
"xander_sd15_lora.safetensors",
|
|
||||||
0.6,
|
|
||||||
0.6
|
|
||||||
]
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"links": [
|
|
||||||
[
|
|
||||||
2,
|
|
||||||
5,
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
3,
|
|
||||||
"LATENT"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
4,
|
|
||||||
6,
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
1,
|
|
||||||
"CONDITIONING"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
6,
|
|
||||||
7,
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
2,
|
|
||||||
"CONDITIONING"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
7,
|
|
||||||
3,
|
|
||||||
0,
|
|
||||||
8,
|
|
||||||
0,
|
|
||||||
"LATENT"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
8,
|
|
||||||
4,
|
|
||||||
2,
|
|
||||||
8,
|
|
||||||
1,
|
|
||||||
"VAE"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
9,
|
|
||||||
8,
|
|
||||||
0,
|
|
||||||
9,
|
|
||||||
0,
|
|
||||||
"IMAGE"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
10,
|
|
||||||
4,
|
|
||||||
0,
|
|
||||||
10,
|
|
||||||
0,
|
|
||||||
"MODEL"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
12,
|
|
||||||
4,
|
|
||||||
1,
|
|
||||||
10,
|
|
||||||
1,
|
|
||||||
"CLIP"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
15,
|
|
||||||
11,
|
|
||||||
0,
|
|
||||||
6,
|
|
||||||
0,
|
|
||||||
"CLIP"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
16,
|
|
||||||
11,
|
|
||||||
0,
|
|
||||||
7,
|
|
||||||
0,
|
|
||||||
"CLIP"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
19,
|
|
||||||
12,
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
0,
|
|
||||||
"MODEL"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
22,
|
|
||||||
10,
|
|
||||||
0,
|
|
||||||
12,
|
|
||||||
0,
|
|
||||||
"*"
|
|
||||||
],
|
|
||||||
[
|
|
||||||
23,
|
|
||||||
10,
|
|
||||||
1,
|
|
||||||
11,
|
|
||||||
0,
|
|
||||||
"*"
|
|
||||||
]
|
|
||||||
],
|
|
||||||
"groups": [],
|
|
||||||
"config": {},
|
|
||||||
"extra": {
|
|
||||||
"ds": {
|
|
||||||
"scale": 0.8264462809917354,
|
|
||||||
"offset": {
|
|
||||||
"0": 513.8734070325743,
|
|
||||||
"1": 351.4824273966635
|
|
||||||
}
|
|
||||||
}
|
|
||||||
},
|
|
||||||
"version": 0.4
|
|
||||||
}
|
|
||||||
@@ -1,28 +0,0 @@
|
|||||||
import numpy as np
|
|
||||||
from PIL import Image
|
|
||||||
import os
|
|
||||||
|
|
||||||
def generate_random_color_image(width, height):
|
|
||||||
"""Generate an image of random color."""
|
|
||||||
color = np.random.randint(0, 256, (3,), dtype=np.uint8)
|
|
||||||
image = np.full((height, width, 3), color, dtype=np.uint8)
|
|
||||||
return Image.fromarray(image)
|
|
||||||
|
|
||||||
def save_images(num_images, width, height, directory):
|
|
||||||
"""Save a specified number of random color images."""
|
|
||||||
for i in range(num_images):
|
|
||||||
image = generate_random_color_image(width, height)
|
|
||||||
image.save(f"{directory}/random_color_image_{i+1}.png")
|
|
||||||
|
|
||||||
# Parameters
|
|
||||||
num_images = 40
|
|
||||||
width = 1024
|
|
||||||
height = 1024
|
|
||||||
directory = "random_images"
|
|
||||||
|
|
||||||
os.makedirs(directory, exist_ok=True)
|
|
||||||
|
|
||||||
# Generate and save images
|
|
||||||
save_images(num_images, width, height, directory)
|
|
||||||
|
|
||||||
print(f'Saved {num_images} random color images to {os.path.abspath(directory)}')
|
|
||||||
@@ -1,179 +0,0 @@
|
|||||||
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"
|
|
||||||
)
|
|
||||||
)
|
|
||||||
@@ -1,4 +1,3 @@
|
|||||||
import fnmatch
|
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
@@ -12,20 +11,18 @@ import torch
|
|||||||
import torch.utils.checkpoint
|
import torch.utils.checkpoint
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
import prodigyopt
|
|
||||||
from typing import Union, Iterable, List, Dict, Tuple, Optional, cast
|
|
||||||
#from diffusers.training_utils import cast_training_params
|
|
||||||
|
|
||||||
from trainer.utils.utils import *
|
from trainer.utils.utils import *
|
||||||
from trainer.checkpoint import save_checkpoint
|
from trainer.checkpoint import save_checkpoint
|
||||||
from trainer.embedding_handler import TokenEmbeddingsHandler
|
from trainer.embedding_handler import TokenEmbeddingsHandler
|
||||||
from trainer.dataset import PreprocessedDataset
|
from trainer.dataset import PreprocessedDataset
|
||||||
from trainer.config import TrainingConfig
|
from trainer.config import TrainingConfig
|
||||||
from trainer.models import print_trainable_parameters, load_models
|
from trainer.models import print_trainable_parameters, load_models
|
||||||
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer
|
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer, compute_token_attention_loss
|
||||||
from trainer.inference import render_images, get_conditioning_signals
|
from trainer.inference import render_images, get_conditioning_signals
|
||||||
from trainer.preprocess import preprocess
|
from trainer.preprocess import preprocess
|
||||||
from trainer.utils.io import make_validation_img_grid
|
from trainer.utils.io import make_validation_img_grid
|
||||||
|
from trainer.ti_cross_attn_loss import init_daam_loss, plot_token_attention_loss
|
||||||
|
|
||||||
from trainer.optimizer import (
|
from trainer.optimizer import (
|
||||||
OptimizerCollection,
|
OptimizerCollection,
|
||||||
get_optimizer_and_peft_models_text_encoder_lora,
|
get_optimizer_and_peft_models_text_encoder_lora,
|
||||||
@@ -34,10 +31,43 @@ from trainer.optimizer import (
|
|||||||
get_unet_optimizer
|
get_unet_optimizer
|
||||||
)
|
)
|
||||||
|
|
||||||
def train(
|
def train(config: TrainingConfig):
|
||||||
config: TrainingConfig,
|
|
||||||
):
|
|
||||||
seed_everything(config.seed)
|
seed_everything(config.seed)
|
||||||
|
weight_dtype = dtype_map[config.weight_type]
|
||||||
|
|
||||||
|
(
|
||||||
|
pipe,
|
||||||
|
tokenizer_one,
|
||||||
|
tokenizer_two,
|
||||||
|
noise_scheduler,
|
||||||
|
text_encoder_one,
|
||||||
|
text_encoder_two,
|
||||||
|
vae,
|
||||||
|
unet,
|
||||||
|
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
|
||||||
|
|
||||||
|
pipe, daam_loss = init_daam_loss(
|
||||||
|
pipeline=pipe
|
||||||
|
)
|
||||||
|
|
||||||
|
config.sd_model_version = sd_model_version
|
||||||
|
config.pretrained_model["version"] = sd_model_version
|
||||||
|
|
||||||
|
if not config.sample_imgs_lora_scale:
|
||||||
|
if config.sd_model_version == "sdxl":
|
||||||
|
config.sample_imgs_lora_scale = 0.75
|
||||||
|
else:
|
||||||
|
config.sample_imgs_lora_scale = 0.85
|
||||||
|
|
||||||
|
if not config.validation_img_size:
|
||||||
|
if config.sd_model_version == "sdxl":
|
||||||
|
config.validation_img_size = 1024
|
||||||
|
else:
|
||||||
|
config.validation_img_size = 768
|
||||||
|
|
||||||
|
print("xxxxxxxxxxxxxxxxxxx")
|
||||||
|
print(config.prompt_modifier)
|
||||||
|
|
||||||
config, input_dir = preprocess(
|
config, input_dir = preprocess(
|
||||||
config,
|
config,
|
||||||
@@ -58,19 +88,6 @@ def train(
|
|||||||
if config.allow_tf32:
|
if config.allow_tf32:
|
||||||
torch.backends.cuda.matmul.allow_tf32 = True
|
torch.backends.cuda.matmul.allow_tf32 = True
|
||||||
|
|
||||||
weight_dtype = dtype_map[config.weight_type]
|
|
||||||
|
|
||||||
(
|
|
||||||
pipe,
|
|
||||||
tokenizer_one,
|
|
||||||
tokenizer_two,
|
|
||||||
noise_scheduler,
|
|
||||||
text_encoder_one,
|
|
||||||
text_encoder_two,
|
|
||||||
vae,
|
|
||||||
unet,
|
|
||||||
) = load_models(config.pretrained_model, config.device, weight_dtype, keep_vae_float32=0)
|
|
||||||
|
|
||||||
# Initialize new tokens for training.
|
# Initialize new tokens for training.
|
||||||
embedding_handler = TokenEmbeddingsHandler(
|
embedding_handler = TokenEmbeddingsHandler(
|
||||||
text_encoders = [text_encoder_one, text_encoder_two],
|
text_encoders = [text_encoder_one, text_encoder_two],
|
||||||
@@ -113,22 +130,26 @@ def train(
|
|||||||
|
|
||||||
|
|
||||||
embedding_handler.make_embeddings_trainable()
|
embedding_handler.make_embeddings_trainable()
|
||||||
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
|
if not config.disable_ti:
|
||||||
text_encoders=text_encoders,
|
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
|
||||||
textual_inversion_lr=config.ti_lr,
|
text_encoders=text_encoders,
|
||||||
textual_inversion_weight_decay=config.ti_weight_decay,
|
textual_inversion_lr=config.ti_lr,
|
||||||
optimizer_name=config.ti_optimizer ## hardcoded
|
textual_inversion_weight_decay=config.ti_weight_decay,
|
||||||
)
|
optimizer_name=config.ti_optimizer ## hardcoded
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
optimizer_ti = None
|
||||||
|
textual_inversion_params = None
|
||||||
|
|
||||||
if not config.is_lora: # This code pathway has not been tested in a long while
|
if not config.is_lora: # This code pathway has not been tested in a long while
|
||||||
print(f"Doing full fine-tuning on the U-Net")
|
print(f"Doing full fine-tuning on the U-Net")
|
||||||
unet.requires_grad_(True)
|
unet.requires_grad_(True)
|
||||||
unet_lora_parameters = None
|
unet_lora_parameters = None
|
||||||
|
optimizer_text_encoder_lora = None
|
||||||
unet_trainable_params = unet.parameters()
|
unet_trainable_params = unet.parameters()
|
||||||
else:
|
else:
|
||||||
# Do lora-training instead.
|
# Do lora-training instead.
|
||||||
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
|
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
|
||||||
|
|
||||||
# target_blocks=["block"] for original IP-Adapter
|
# target_blocks=["block"] for original IP-Adapter
|
||||||
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
|
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
|
||||||
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
|
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
|
||||||
@@ -142,15 +163,18 @@ def train(
|
|||||||
pipe=pipe
|
pipe=pipe
|
||||||
)
|
)
|
||||||
|
|
||||||
optimizer_unet = get_unet_optimizer(
|
if config.unet_lr > 0.0:
|
||||||
prodigy_d_coef=config.prodigy_d_coef,
|
optimizer_unet = get_unet_optimizer(
|
||||||
prodigy_growth_factor=config.unet_prodigy_growth_factor,
|
prodigy_d_coef=config.prodigy_d_coef,
|
||||||
lora_weight_decay=config.lora_weight_decay,
|
prodigy_growth_factor=config.unet_prodigy_growth_factor,
|
||||||
use_dora=config.use_dora,
|
lora_weight_decay=config.lora_weight_decay,
|
||||||
unet_trainable_params=unet_trainable_params,
|
use_dora=config.use_dora,
|
||||||
optimizer_name=config.unet_optimizer_type
|
unet_trainable_params=unet_trainable_params,
|
||||||
)
|
optimizer_name=config.unet_optimizer_type
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
optimizer_unet = None
|
||||||
|
|
||||||
print_trainable_parameters(unet, model_name = 'unet')
|
print_trainable_parameters(unet, model_name = 'unet')
|
||||||
for i, text_encoder in enumerate(text_encoders):
|
for i, text_encoder in enumerate(text_encoders):
|
||||||
if text_encoder is not None:
|
if text_encoder is not None:
|
||||||
@@ -161,31 +185,26 @@ def train(
|
|||||||
pipe,
|
pipe,
|
||||||
vae.float(),
|
vae.float(),
|
||||||
size = config.train_img_size,
|
size = config.train_img_size,
|
||||||
do_cache=config.do_cache,
|
|
||||||
substitute_caption_map=config.token_dict,
|
substitute_caption_map=config.token_dict,
|
||||||
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
|
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
|
||||||
train_batch_size=config.train_batch_size
|
train_batch_size=config.train_batch_size
|
||||||
)
|
)
|
||||||
# offload the vae to cpu:
|
print("Final training captions:")
|
||||||
|
print(train_dataset.captions[:40])
|
||||||
|
|
||||||
|
# offload the vae to cpu and release memory:
|
||||||
vae = vae.to('cpu')
|
vae = vae.to('cpu')
|
||||||
gc.collect()
|
gc.collect()
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
|
|
||||||
train_dataloader = torch.utils.data.DataLoader(
|
train_dataloader = torch.utils.data.DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
batch_size=config.train_batch_size,
|
batch_size=config.train_batch_size,
|
||||||
shuffle=True,
|
shuffle=True,
|
||||||
num_workers=config.dataloader_num_workers,
|
num_workers=config.dataloader_num_workers
|
||||||
)
|
)
|
||||||
|
|
||||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps)
|
config.num_train_epochs = int(math.ceil(config.max_train_steps / len(train_dataloader)))
|
||||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
|
|
||||||
|
|
||||||
if config.max_train_steps is None:
|
|
||||||
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
|
|
||||||
|
|
||||||
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
|
|
||||||
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
|
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
|
||||||
|
|
||||||
print(f"--- Num samples = {len(train_dataset)}")
|
print(f"--- Num samples = {len(train_dataset)}")
|
||||||
@@ -194,7 +213,7 @@ def train(
|
|||||||
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
|
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
|
||||||
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
|
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
|
||||||
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
|
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
|
||||||
print(f"--- Total optimization steps = {config.max_train_steps}\n")
|
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
|
||||||
|
|
||||||
global_step = 0
|
global_step = 0
|
||||||
last_save_step = 0
|
last_save_step = 0
|
||||||
@@ -208,20 +227,18 @@ def train(
|
|||||||
# Data tracking inits:
|
# Data tracking inits:
|
||||||
start_time, images_done = time.time(), 0
|
start_time, images_done = time.time(), 0
|
||||||
prompt_embeds_norms = {'main':[], 'reg':[]}
|
prompt_embeds_norms = {'main':[], 'reg':[]}
|
||||||
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': []}
|
losses = {'img_loss': [], 'tot_loss': [], 'covariance_tok_reg_loss': [], 'concept_description_loss': [], 'token_std_loss': [], 'token_attention_loss': []}
|
||||||
grad_norms, token_stds = {'unet': []}, {}
|
grad_norms, token_stds = {'unet': []}, {}
|
||||||
for i in range(len(text_encoders)):
|
for i in range(len(text_encoders)):
|
||||||
grad_norms[f'text_encoder_{i}'] = []
|
grad_norms[f'text_encoder_{i}'] = []
|
||||||
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
|
token_stds[f'text_encoder_{i}'] = {j: [] for j in range(config.n_tokens)}
|
||||||
|
|
||||||
|
# default values for cold (starting) optimizer lr:
|
||||||
|
base_unet_lr = 2.0e-4 if (config.is_lora and config.disable_ti) else 5.0e-5
|
||||||
|
|
||||||
# default value of cold (pre-warmup) optimizer lr:
|
if not config.is_lora:
|
||||||
if config.sd_model_version == "sdxl":
|
base_unet_lr = 1.0e-5
|
||||||
# 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
|
|
||||||
|
|
||||||
#######################################################################################################
|
#######################################################################################################
|
||||||
|
|
||||||
"""
|
"""
|
||||||
@@ -235,7 +252,8 @@ def train(
|
|||||||
)
|
)
|
||||||
optimizers = optimizer_collection.optimizers
|
optimizers = optimizer_collection.optimizers
|
||||||
|
|
||||||
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
|
if config.debug:
|
||||||
|
embedding_handler.visualize_random_token_embeddings(os.path.join(config.output_dir, 'ti_embeddings'), n = 10)
|
||||||
|
|
||||||
for epoch in range(config.num_train_epochs):
|
for epoch in range(config.num_train_epochs):
|
||||||
if config.aspect_ratio_bucketing:
|
if config.aspect_ratio_bucketing:
|
||||||
@@ -248,15 +266,12 @@ def train(
|
|||||||
completion_f = finegrained_epoch / config.num_train_epochs
|
completion_f = finegrained_epoch / config.num_train_epochs
|
||||||
|
|
||||||
# param_groups[1] goes from ti_lr to 0.0 over the course of training
|
# param_groups[1] goes from ti_lr to 0.0 over the course of training
|
||||||
if config.ti_optimizer != "prodigy": # Update ti_learning rate gradually:
|
if config.ti_optimizer != "prodigy" and optimizers['textual_inversion'] is not None:
|
||||||
if optimizers['textual_inversion'] is not None:
|
# Apply the exponential learning rate
|
||||||
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0
|
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 1.7
|
||||||
# warmup the ti-lr:
|
# Apply freezing condition
|
||||||
if config.ti_lr_warmup_steps > 0:
|
if completion_f > config.freeze_ti_after_completion_f:
|
||||||
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0)
|
optimizers['textual_inversion'].param_groups[0]['lr'] = 0.0
|
||||||
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
|
|
||||||
if config.freeze_ti_after_completion_f <= completion_f:
|
|
||||||
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
|
|
||||||
|
|
||||||
if optimizers['text_encoders'] is not None:
|
if optimizers['text_encoders'] is not None:
|
||||||
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
|
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
|
||||||
@@ -268,16 +283,26 @@ def train(
|
|||||||
|
|
||||||
if optimizers['unet'] is not None:
|
if optimizers['unet'] is not None:
|
||||||
# Calculate the exponential factor
|
# Calculate the exponential factor
|
||||||
exp_factor = (config.unet_lr / base_lr) ** (global_step / config.unet_lr_warmup_steps)
|
exp_factor = (config.unet_lr / base_unet_lr) ** (global_step / config.unet_lr_warmup_steps)
|
||||||
# Apply the exponential learning rate
|
# Apply the exponential learning rate
|
||||||
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
|
optimizers['unet'].param_groups[0]['lr'] = base_unet_lr * exp_factor
|
||||||
|
|
||||||
|
if completion_f < config.freeze_unet_before_completion_f:
|
||||||
|
optimizers['unet'].param_groups[0]['lr'] = 0.0
|
||||||
|
|
||||||
if not config.aspect_ratio_bucketing:
|
if not config.aspect_ratio_bucketing:
|
||||||
captions, vae_latent, mask = batch
|
captions, vae_latent, mask = batch
|
||||||
else:
|
else:
|
||||||
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
|
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
|
||||||
|
|
||||||
|
mask = mask.to(config.device)
|
||||||
|
|
||||||
captions = list(captions)
|
captions = list(captions)
|
||||||
|
if config.caption_dropout > 0.0:
|
||||||
|
for i in range(len(captions)):
|
||||||
|
if np.random.rand() < config.caption_dropout:
|
||||||
|
captions[i] = config.token_dict["TOK"]
|
||||||
|
|
||||||
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
|
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
|
||||||
config, pipe, captions
|
config, pipe, captions
|
||||||
)
|
)
|
||||||
@@ -309,18 +334,23 @@ def train(
|
|||||||
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
|
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
|
||||||
return_dict=False,
|
return_dict=False,
|
||||||
)[0]
|
)[0]
|
||||||
|
|
||||||
# Compute the loss:
|
# Compute the loss:
|
||||||
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
|
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
|
||||||
losses['img_loss'].append(loss.item())
|
losses['img_loss'].append(loss.item())
|
||||||
|
|
||||||
|
if not config.disable_ti:
|
||||||
|
token_attention_loss = compute_token_attention_loss(pipe, embedding_handler, captions, mask, daam_loss)
|
||||||
|
losses['token_attention_loss'].append(token_attention_loss.item())
|
||||||
|
loss = loss + config.token_attention_loss_w * token_attention_loss
|
||||||
|
|
||||||
if config.training_attributes["gpt_description"] and config.debug:
|
if config.training_attributes["gpt_description"] and config.debug:
|
||||||
concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe)
|
concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe)
|
||||||
# Dont apply this loss, just plot it for now:
|
# Dont apply this loss, just plot it for now:
|
||||||
loss += 0.0 * concept_description_loss
|
loss += 0.0 * concept_description_loss
|
||||||
losses['concept_description_loss'].append(concept_description_loss.item())
|
losses['concept_description_loss'].append(concept_description_loss.item())
|
||||||
|
|
||||||
if config.l1_penalty > 0.0:
|
if config.l1_penalty > 0.0 and unet_lora_parameters:
|
||||||
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
|
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
|
||||||
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
|
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
|
||||||
loss += config.l1_penalty * l1_norm
|
loss += config.l1_penalty * l1_norm
|
||||||
@@ -349,12 +379,6 @@ def train(
|
|||||||
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
|
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
|
||||||
|
|
||||||
optimizer_collection.step()
|
optimizer_collection.step()
|
||||||
|
|
||||||
# after every optimizer step, we do some manual intervention of the embeddings to regularize them:
|
|
||||||
if optimizer_collection.get_lr('textual_inversion') > 0.0:
|
|
||||||
#embedding_handler.fix_embedding_std(config.off_ratio_power)
|
|
||||||
pass
|
|
||||||
|
|
||||||
optimizer_collection.zero_grad()
|
optimizer_collection.zero_grad()
|
||||||
|
|
||||||
#############################################################################################################
|
#############################################################################################################
|
||||||
@@ -368,9 +392,14 @@ def train(
|
|||||||
for std_i, std in enumerate(embedding_stds):
|
for std_i, std in enumerate(embedding_stds):
|
||||||
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
|
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
|
||||||
|
|
||||||
|
if global_step % 50 == 0 and not config.disable_ti and config.debug:
|
||||||
|
img_ratio = config.train_img_size[0] / config.train_img_size[1]
|
||||||
|
plot_token_attention_loss(config.output_dir, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio)
|
||||||
|
|
||||||
# Print some statistics:
|
# Print some statistics:
|
||||||
if config.debug and (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > -1:
|
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)): #and global_step > 0:
|
||||||
|
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
|
||||||
|
|
||||||
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
|
||||||
os.makedirs(output_save_dir, exist_ok=True)
|
os.makedirs(output_save_dir, exist_ok=True)
|
||||||
config.save_as_json(
|
config.save_as_json(
|
||||||
@@ -390,27 +419,17 @@ def train(
|
|||||||
)
|
)
|
||||||
last_save_step = global_step
|
last_save_step = global_step
|
||||||
|
|
||||||
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings()
|
if config.debug:
|
||||||
for idx, text_encoder in enumerate(text_encoders):
|
embedding_handler.print_token_info()
|
||||||
if text_encoder is None:
|
if config.is_lora: # plotting this hist for full unet parameters can run OOM
|
||||||
continue
|
plot_torch_hist(unet_lora_parameters, global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
|
||||||
n = len(token_embeddings[f'txt_encoder_{idx}'])
|
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
|
||||||
for i in range(n):
|
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
|
||||||
token = trainable_tokens[f'txt_encoder_{idx}'][i]
|
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
|
||||||
# Strip any backslashes from the token name:
|
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
|
||||||
token = token.replace("/", "_")
|
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
|
||||||
embedding = token_embeddings[f'txt_encoder_{idx}'][i]
|
#plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
|
||||||
plot_torch_hist(embedding, global_step, os.path.join(config.output_dir, 'ti_embeddings') , f"enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'red')
|
|
||||||
|
|
||||||
embedding_handler.print_token_info()
|
|
||||||
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
|
|
||||||
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
|
|
||||||
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
|
|
||||||
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
|
|
||||||
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
|
|
||||||
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
|
|
||||||
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
|
|
||||||
|
|
||||||
validation_prompts = render_images(
|
validation_prompts = render_images(
|
||||||
pipe = pipe,
|
pipe = pipe,
|
||||||
render_size = config.validation_img_size,
|
render_size = config.validation_img_size,
|
||||||
@@ -420,6 +439,8 @@ def train(
|
|||||||
is_lora = config.is_lora,
|
is_lora = config.is_lora,
|
||||||
pretrained_model = config.pretrained_model,
|
pretrained_model = config.pretrained_model,
|
||||||
lora_scale = config.sample_imgs_lora_scale,
|
lora_scale = config.sample_imgs_lora_scale,
|
||||||
|
disable_ti = config.disable_ti,
|
||||||
|
prompt_modifier = config.prompt_modifier,
|
||||||
n_imgs = config.n_sample_imgs,
|
n_imgs = config.n_sample_imgs,
|
||||||
device = config.device,
|
device = config.device,
|
||||||
checkpoint_folder = None
|
checkpoint_folder = None
|
||||||
@@ -433,14 +454,13 @@ def train(
|
|||||||
images_done += config.train_batch_size
|
images_done += config.train_batch_size
|
||||||
global_step += 1
|
global_step += 1
|
||||||
|
|
||||||
if global_step % (config.max_train_steps//20) == 0:
|
if global_step % (config.max_train_steps//100) == 0:
|
||||||
progress = (global_step / config.max_train_steps) + 0.05
|
progress = (global_step / config.max_train_steps) + 0.05
|
||||||
print_system_info()
|
#print_system_info()
|
||||||
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
|
|
||||||
yield np.min((progress, 1.0))
|
yield np.min((progress, 1.0))
|
||||||
|
|
||||||
if global_step > config.max_train_steps:
|
if global_step > config.max_train_steps:
|
||||||
print("Reached max steps, stopping training!")
|
print("Reached max steps, stopping training!", flush = True)
|
||||||
break
|
break
|
||||||
|
|
||||||
# final_save
|
# final_save
|
||||||
@@ -471,8 +491,7 @@ def train(
|
|||||||
pretrained_model_version=config.pretrained_model["version"]
|
pretrained_model_version=config.pretrained_model["version"]
|
||||||
)
|
)
|
||||||
|
|
||||||
print("Running final inference round...")
|
if config.debug and 0:
|
||||||
if config.debug:
|
|
||||||
# Reload the entire pipe from disk + LoRa:
|
# Reload the entire pipe from disk + LoRa:
|
||||||
pipe_to_use = None
|
pipe_to_use = None
|
||||||
checkpoint_folder = output_save_dir
|
checkpoint_folder = output_save_dir
|
||||||
@@ -502,6 +521,8 @@ def train(
|
|||||||
is_lora=config.is_lora,
|
is_lora=config.is_lora,
|
||||||
pretrained_model=config.pretrained_model,
|
pretrained_model=config.pretrained_model,
|
||||||
lora_scale=config.sample_imgs_lora_scale,
|
lora_scale=config.sample_imgs_lora_scale,
|
||||||
|
disable_ti = config.disable_ti,
|
||||||
|
prompt_modifier = config.prompt_modifier,
|
||||||
n_imgs = config.n_sample_imgs,
|
n_imgs = config.n_sample_imgs,
|
||||||
n_steps = 30,
|
n_steps = 30,
|
||||||
device = config.device,
|
device = config.device,
|
||||||
@@ -511,13 +532,6 @@ def train(
|
|||||||
img_grid_path = make_validation_img_grid(output_save_dir)
|
img_grid_path = make_validation_img_grid(output_save_dir)
|
||||||
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
|
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
|
||||||
|
|
||||||
# Remove unneeded checkpoints if they exist in the output directory:
|
|
||||||
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
|
|
||||||
for file in to_remove:
|
|
||||||
file_path = os.path.join(output_save_dir, file)
|
|
||||||
if os.path.exists(file_path):
|
|
||||||
os.remove(file_path)
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
print(f"Skipping final save, {output_save_dir} already exists")
|
print(f"Skipping final save, {output_save_dir} already exists")
|
||||||
|
|
||||||
@@ -531,6 +545,8 @@ def train(
|
|||||||
config.job_time = time.time() - config.start_time
|
config.job_time = time.time() - config.start_time
|
||||||
config.training_attributes["validation_prompts"] = validation_prompts
|
config.training_attributes["validation_prompts"] = validation_prompts
|
||||||
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
|
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
|
||||||
|
print("Training job complete, saving outputs...", flush = True)
|
||||||
|
print("------------------------------------------")
|
||||||
|
|
||||||
return config, output_save_dir
|
return config, output_save_dir
|
||||||
|
|
||||||
@@ -541,7 +557,12 @@ if __name__ == "__main__":
|
|||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
config = TrainingConfig.from_json(file_path=args.config_filename)
|
config = TrainingConfig.from_json(file_path=args.config_filename)
|
||||||
|
|
||||||
|
print("Starting new LoRa training run with config:")
|
||||||
|
print(config)
|
||||||
|
print("------------------------------------------")
|
||||||
|
|
||||||
for progress in train(config=config):
|
for progress in train(config=config):
|
||||||
print(f"Progress: {(100*progress):.2f}%", end="\r")
|
print(f"Progress: {(100*progress):.2f}%", end="\r")
|
||||||
|
|
||||||
print("Training done :)")
|
print("Training done :)")
|
||||||
|
|||||||
-2027
File diff suppressed because it is too large
Load Diff
@@ -1,22 +1,17 @@
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
import shutil
|
|
||||||
import tarfile
|
import tarfile
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
import random
|
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
from PIL import Image
|
||||||
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
from main import train
|
from main import train
|
||||||
|
from trainer.config import TrainingConfig, model_paths
|
||||||
from trainer.preprocess import preprocess
|
|
||||||
from trainer.models import pretrained_models
|
|
||||||
from trainer.config import TrainingConfig
|
|
||||||
from trainer.utils.io import clean_filename
|
from trainer.utils.io import clean_filename
|
||||||
from trainer.utils.utils import seed_everything
|
|
||||||
|
import folder_paths
|
||||||
|
import comfy.utils
|
||||||
|
|
||||||
class Eden_LoRa_trainer:
|
class Eden_LoRa_trainer:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -24,92 +19,130 @@ class Eden_LoRa_trainer:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"training_images_folder_path": ("STRING", {"default": "."}),
|
"training_images_folder_path": ("STRING", {"default": "."}),
|
||||||
"lora_name": ("STRING", {"default": ""}),
|
"mode": (["style", "face", "object"], {"default": "style"}),
|
||||||
"sd_model_version": (["sdxl", "sd15"], ),
|
"lora_name": ("STRING", {"default": "Eden_Token_LoRa"}),
|
||||||
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||||
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
|
"training_resolution": ("INT", {"default": 512, "min": 256, "max": 1024}),
|
||||||
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
|
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
|
||||||
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
|
"max_train_steps": ("INT", {"default": 300, "min": 10, "max": 10000}),
|
||||||
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
|
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0, "max": 0.005, "step": 0.0001}),
|
||||||
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
|
"unet_lr": ("FLOAT", {"default": 0.0005, "min": 0.0, "max": 0.005, "step": 0.0001}),
|
||||||
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
|
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
|
||||||
"use_dora": ("BOOLEAN", {"default": False}),
|
"disable_ti": ("BOOLEAN", {"default": False}),
|
||||||
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}),
|
"n_tokens": ("INT", {"default": 3, "min": 1, "max": 5}),
|
||||||
|
"save_checkpoint_every_n_steps": ("INT", {"default": 200, "min": 10, "max": 10000}),
|
||||||
|
"n_sample_imgs": ("INT", {"default": 4, "min": 2, "max": 10}),
|
||||||
|
"sample_imgs_lora_scale": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.25}),
|
||||||
|
"plot_training_graphs_on_disk": ("BOOLEAN", {"default": False}),
|
||||||
|
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
CATEGORY = "Eden 🌱"
|
CATEGORY = "Eden 🌱"
|
||||||
RETURN_TYPES = ("STRING",)
|
RETURN_TYPES = ("IMAGE", "STRING", "STRING", "STRING")
|
||||||
|
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
|
||||||
FUNCTION = "train_lora"
|
FUNCTION = "train_lora"
|
||||||
|
|
||||||
def train_lora(self, training_images_folder_path,
|
def train_lora(self,
|
||||||
name = lora_name,
|
training_images_folder_path,
|
||||||
concept_mode = "style",
|
ckpt_name,
|
||||||
sd_model_version = "sdxl",
|
lora_name,
|
||||||
seed = 0,
|
mode,
|
||||||
resolution = 521,
|
training_resolution,
|
||||||
train_batch_size = 4,
|
train_batch_size,
|
||||||
max_train_steps = 400,
|
max_train_steps ,
|
||||||
ti_lr = 0.001,
|
ti_lr,
|
||||||
unet_lr = 0.001,
|
unet_lr,
|
||||||
lora_rank = 16,
|
lora_rank,
|
||||||
use_dora = False,
|
disable_ti,
|
||||||
n_tokens = 2
|
n_tokens,
|
||||||
|
plot_training_graphs_on_disk,
|
||||||
|
save_checkpoint_every_n_steps,
|
||||||
|
n_sample_imgs,
|
||||||
|
sample_imgs_lora_scale,
|
||||||
|
seed,
|
||||||
):
|
):
|
||||||
|
|
||||||
print("Starting new training job...")
|
print("Starting new training job...")
|
||||||
|
|
||||||
|
# Overwrite hardcoded paths to point to comfyUI folders:
|
||||||
|
model_paths.set_path("CLIP", os.path.join(folder_paths.models_dir, "clipseg"))
|
||||||
|
model_paths.set_path("FLORENCE", os.path.join(folder_paths.models_dir, "LLM"))
|
||||||
|
model_paths.set_path("BLIP", os.path.join(folder_paths.models_dir, "blip"))
|
||||||
|
model_paths.set_path("SR", os.path.join(folder_paths.models_dir, "upscale_models"))
|
||||||
|
model_paths.set_path("SD", os.path.join(folder_paths.models_dir, "checkpoints"))
|
||||||
|
|
||||||
|
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||||
|
|
||||||
config = TrainingConfig(
|
config = TrainingConfig(
|
||||||
name="test",
|
name=lora_name,
|
||||||
|
output_dir="output",
|
||||||
lora_training_urls=training_images_folder_path,
|
lora_training_urls=training_images_folder_path,
|
||||||
concept_mode=concept_mode,
|
concept_mode=mode,
|
||||||
sd_model_version=sd_model_version,
|
ckpt_path=ckpt_path,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
resolution=resolution,
|
resolution=training_resolution,
|
||||||
train_batch_size=train_batch_size,
|
train_batch_size=train_batch_size,
|
||||||
max_train_steps=max_train_steps,
|
max_train_steps=max_train_steps,
|
||||||
checkpointing_steps=10000,
|
checkpointing_steps=save_checkpoint_every_n_steps,
|
||||||
|
n_sample_imgs=(n_sample_imgs//2) * 2,
|
||||||
|
sample_imgs_lora_scale=sample_imgs_lora_scale,
|
||||||
ti_lr=ti_lr,
|
ti_lr=ti_lr,
|
||||||
unet_lr=unet_lr,
|
unet_lr=unet_lr,
|
||||||
lora_rank=lora_rank,
|
lora_rank=lora_rank,
|
||||||
use_dora=use_dora,
|
use_dora=False,
|
||||||
caption_model="blip",
|
caption_model="blip",
|
||||||
|
disable_ti=disable_ti,
|
||||||
n_tokens=n_tokens,
|
n_tokens=n_tokens,
|
||||||
verbose=True,
|
verbose=True,
|
||||||
debug=True,
|
debug=plot_training_graphs_on_disk,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(100)
|
||||||
|
|
||||||
with torch.inference_mode(False):
|
with torch.inference_mode(False):
|
||||||
train_generator = train(config=config)
|
train_generator = train(config=config)
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
progress_f = next(train_generator)
|
progress_f = next(train_generator)
|
||||||
|
pbar.update_absolute(progress_f * 100)
|
||||||
except StopIteration as e:
|
except StopIteration as e:
|
||||||
config, output_save_dir = e.value # Capture the return value
|
config, output_save_dir = e.value # Capture the return value
|
||||||
break
|
break
|
||||||
|
|
||||||
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
|
|
||||||
out_path = f"{clean_filename(lora_name)}_eden_concept_lora_{int(time.time())}.tar"
|
|
||||||
directory = cogPath(output_save_dir)
|
|
||||||
|
|
||||||
with tarfile.open(out_path, "w") as tar:
|
|
||||||
print("Adding files to tar...")
|
|
||||||
for file_path in directory.rglob("*"):
|
|
||||||
print(file_path)
|
|
||||||
arcname = file_path.relative_to(directory)
|
|
||||||
tar.add(file_path, arcname=arcname)
|
|
||||||
|
|
||||||
# Add instructions README:
|
|
||||||
tar.add("instructions_README.md", arcname="README.md")
|
|
||||||
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
|
|
||||||
if sd_model_version == "sd15":
|
|
||||||
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
|
|
||||||
|
|
||||||
attributes = {}
|
attributes = {}
|
||||||
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
||||||
attributes['job_time_seconds'] = config.job_time
|
attributes['job_time_seconds'] = config.job_time
|
||||||
|
|
||||||
print(f"LORA training finished in {config.job_time:.1f} seconds")
|
print(f"LORA training node finished in {config.job_time:.1f} seconds")
|
||||||
print(f"Returning {out_path}")
|
print("---------- Made with love by Eden.art 🌱 ----------")
|
||||||
|
|
||||||
|
# safetensors paths:
|
||||||
|
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
|
||||||
|
|
||||||
return (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)
|
||||||
+35
-18
@@ -14,7 +14,6 @@ from main import train
|
|||||||
from typing import Iterator, Optional
|
from typing import Iterator, Optional
|
||||||
|
|
||||||
from trainer.preprocess import preprocess
|
from trainer.preprocess import preprocess
|
||||||
from trainer.models import pretrained_models
|
|
||||||
from trainer.config import TrainingConfig
|
from trainer.config import TrainingConfig
|
||||||
from trainer.utils.io import clean_filename
|
from trainer.utils.io import clean_filename
|
||||||
from trainer.utils.utils import seed_everything
|
from trainer.utils.utils import seed_everything
|
||||||
@@ -66,19 +65,19 @@ class Predictor(BasePredictor):
|
|||||||
),
|
),
|
||||||
max_train_steps: int = Input(
|
max_train_steps: int = Input(
|
||||||
description="Number of training steps. Increasing this usually leads to overfitting, only viable if you have > 100 training imgs. For faces you may want to reduce to eg 300",
|
description="Number of training steps. Increasing this usually leads to overfitting, only viable if you have > 100 training imgs. For faces you may want to reduce to eg 300",
|
||||||
default=400
|
default=300
|
||||||
|
),
|
||||||
|
checkpointing_steps: int = Input(
|
||||||
|
description="Save a checkpoint every n steps (The final checkpoint will always be saved)",
|
||||||
|
default=10000
|
||||||
),
|
),
|
||||||
resolution: int = Input(
|
resolution: int = Input(
|
||||||
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
|
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 768",
|
||||||
default=512
|
default=512
|
||||||
),
|
),
|
||||||
train_batch_size: int = Input(
|
|
||||||
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
|
|
||||||
default=4
|
|
||||||
),
|
|
||||||
unet_lr: float = Input(
|
unet_lr: float = Input(
|
||||||
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
|
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
|
||||||
default=0.001
|
default=0.0003
|
||||||
),
|
),
|
||||||
ti_lr: float = Input(
|
ti_lr: float = Input(
|
||||||
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
|
description="Learning rate for training textual inversion embeddings. Don't alter unless you know what you're doing.",
|
||||||
@@ -88,13 +87,25 @@ class Predictor(BasePredictor):
|
|||||||
description="Rank of LoRA embeddings for the unet.",
|
description="Rank of LoRA embeddings for the unet.",
|
||||||
default=16
|
default=16
|
||||||
),
|
),
|
||||||
use_dora: bool = Input(
|
|
||||||
description="Use Dora instead of LoRa",
|
|
||||||
default=False,
|
|
||||||
),
|
|
||||||
n_tokens: int = Input(
|
n_tokens: int = Input(
|
||||||
description="How many new tokens to train (highly recommended to leave this at 2)",
|
description="How many new tokens to train (highly recommended to leave this at 2)",
|
||||||
ge=1, le=3, default=2
|
ge=1, le=4, default=3
|
||||||
|
),
|
||||||
|
train_batch_size: int = Input(
|
||||||
|
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
|
||||||
|
default=4
|
||||||
|
),
|
||||||
|
n_sample_imgs: int = Input(
|
||||||
|
description="Number of sample images in validation grid",
|
||||||
|
default=4
|
||||||
|
),
|
||||||
|
validation_img_size: int = Input(
|
||||||
|
description="Resolution of sample images in validation grid",
|
||||||
|
default=1024
|
||||||
|
),
|
||||||
|
sample_imgs_lora_scale: float = Input(
|
||||||
|
description="Scale factor for LoRa when generating sample images. If not provided, will be set automatically",
|
||||||
|
default=None
|
||||||
),
|
),
|
||||||
seed: int = Input(
|
seed: int = Input(
|
||||||
description="Random seed for reproducible training. Leave empty to use a random seed",
|
description="Random seed for reproducible training. Leave empty to use a random seed",
|
||||||
@@ -124,13 +135,15 @@ class Predictor(BasePredictor):
|
|||||||
sd_model_version=sd_model_version,
|
sd_model_version=sd_model_version,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
resolution=resolution,
|
resolution=resolution,
|
||||||
|
validation_img_size=[validation_img_size, validation_img_size],
|
||||||
|
sample_imgs_lora_scale=sample_imgs_lora_scale,
|
||||||
train_batch_size=train_batch_size,
|
train_batch_size=train_batch_size,
|
||||||
max_train_steps=max_train_steps,
|
max_train_steps=max_train_steps,
|
||||||
checkpointing_steps=10000,
|
checkpointing_steps=checkpointing_steps,
|
||||||
|
n_sample_imgs=n_sample_imgs,
|
||||||
ti_lr=ti_lr,
|
ti_lr=ti_lr,
|
||||||
unet_lr=unet_lr,
|
unet_lr=unet_lr,
|
||||||
lora_rank=lora_rank,
|
lora_rank=lora_rank,
|
||||||
use_dora=use_dora,
|
|
||||||
caption_model="blip",
|
caption_model="blip",
|
||||||
n_tokens=n_tokens,
|
n_tokens=n_tokens,
|
||||||
verbose=True,
|
verbose=True,
|
||||||
@@ -162,9 +175,13 @@ class Predictor(BasePredictor):
|
|||||||
|
|
||||||
# Add instructions README:
|
# Add instructions README:
|
||||||
tar.add("instructions_README.md", arcname="README.md")
|
tar.add("instructions_README.md", arcname="README.md")
|
||||||
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
|
comfy_workflows_path = "ComfyUI_workflows"
|
||||||
if sd_model_version == "sd15":
|
if os.path.exists(comfy_workflows_path) and os.path.isdir(comfy_workflows_path):
|
||||||
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
|
for root, dirs, files in os.walk(comfy_workflows_path):
|
||||||
|
for file in files:
|
||||||
|
file_path = os.path.join(root, file)
|
||||||
|
arcname = os.path.relpath(file_path, os.path.dirname(comfy_workflows_path))
|
||||||
|
tar.add(file_path, arcname=arcname)
|
||||||
|
|
||||||
attributes = {}
|
attributes = {}
|
||||||
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
[project]
|
||||||
|
name = "Eden LoRa Trainer"
|
||||||
|
description = "Highly optimized and fast LoRa trainer for SD15 and SDXL models. Does LoRa training, textual inversion and full finetuning!"
|
||||||
|
version = "1.0.0"
|
||||||
|
license = { text = "Other" }
|
||||||
|
dependencies = ["torch==2.1.0", "torchaudio==2.1.0", "torchvision==0.16.0", "transformers==4.38.0", "diffusers==0.29.2", "tokenizers==0.15.2", "huggingface-hub==0.23.2", "ujson==5.10.0", "scipy==1.14.0", "peft==0.10.0", "invisible-watermark==0.2.0", "pandas==2.2.1", "numpy==1.26.4", "opencv-python==4.10.0.84", "mediapipe==0.10.14", "openai==1.35.13", "python-dotenv==1.0.1", "prodigyopt==1.0", "omegaconf==2.3.0", "ujson==5.10.0", "bitsandbytes==0.43.1", "setuptools==70.3.0", "torchtyping==0.1.5", "einops==0.8.0", "timm==1.0.8"]
|
||||||
|
|
||||||
|
[project.urls]
|
||||||
|
Repository = "https://github.com/edenartlab/sd-lora-trainer"
|
||||||
|
Website = "https://www.eden.art/"
|
||||||
|
# Used by Comfy Registry https://comfyregistry.org
|
||||||
|
|
||||||
|
[tool.comfy]
|
||||||
|
PublisherId = "eden"
|
||||||
|
DisplayName = "Eden LoRa Trainer"
|
||||||
|
Icon = ""
|
||||||
+23
-14
@@ -1,17 +1,26 @@
|
|||||||
torch>=2.1.0
|
torch>=2.1.0
|
||||||
|
torchaudio>=2.1.0
|
||||||
torchvision>=0.16.0
|
torchvision>=0.16.0
|
||||||
transformers>=4.38.1
|
transformers>=4.38.0
|
||||||
diffusers>=0.27.2
|
diffusers>=0.29.2
|
||||||
ujson>=5.9.0
|
tokenizers>=0.15.2
|
||||||
scipy>=1.12.0
|
huggingface-hub==0.23.2
|
||||||
peft>=0.10.0
|
ujson==5.10.0
|
||||||
invisible-watermark>=0.2.0
|
scipy==1.14.0
|
||||||
|
peft==0.10.0
|
||||||
|
invisible-watermark==0.2.0
|
||||||
pandas==2.2.1
|
pandas==2.2.1
|
||||||
numpy>=1.26.4
|
numpy==1.26.4
|
||||||
opencv-python>=4.1.0.25
|
opencv-python==4.10.0.84
|
||||||
mediapipe>=0.10.11
|
mediapipe==0.10.14
|
||||||
openai>=1.14.0
|
openai==1.35.13
|
||||||
python-dotenv
|
python-dotenv==1.0.1
|
||||||
prodigyopt
|
prodigyopt==1.0
|
||||||
omegaconf
|
omegaconf==2.3.0
|
||||||
ujson
|
ujson==5.10.0
|
||||||
|
bitsandbytes==0.43.1
|
||||||
|
setuptools==70.3.0
|
||||||
|
torchtyping==0.1.5
|
||||||
|
einops==0.8.0
|
||||||
|
timm==1.0.8
|
||||||
|
triton<3.2.0
|
||||||
|
|||||||
@@ -1,10 +1,8 @@
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
Faces:
|
Faces:
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
|
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
|
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
|
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
|
|
||||||
|
|
||||||
Objects:
|
Objects:
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
|
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
|
||||||
@@ -14,7 +12,12 @@ https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets
|
|||||||
|
|
||||||
Styles:
|
Styles:
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
|
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
|
||||||
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
|
https://edenartlab-lfs.s3.amazonaws.com/datasets/clipx.zip
|
||||||
|
|
||||||
|
/home/rednax/Documents/datasets/good_styles/eden_crystals
|
||||||
|
/home/rednax/Documents/datasets/good_styles/does2
|
||||||
|
/home/rednax/Documents/datasets/good_styles/beeple
|
||||||
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -35,12 +38,12 @@ def hamming_distance(dict1, dict2):
|
|||||||
#######################################################################################
|
#######################################################################################
|
||||||
|
|
||||||
# Setup the base experiment config:
|
# Setup the base experiment config:
|
||||||
exp_name = "grimes"
|
exp_name = "ygor_sd15"
|
||||||
caption_prefix = ""
|
caption_prefix = ""
|
||||||
mask_target_prompts = ""
|
mask_target_prompts = ""
|
||||||
n_exp = 200 # how many random experiment settings to generate
|
n_exp = 100 # how many random experiment settings to generate
|
||||||
min_hamming_distance = 3 # min_n_params that have to be different from any previous experiment to be scheduled
|
min_hamming_distance = 2 # min_n_params that have to be different from any previous experiment to be scheduled
|
||||||
|
nohup = False
|
||||||
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
|
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
|
||||||
|
|
||||||
# Define training hyperparameters and their possible values
|
# Define training hyperparameters and their possible values
|
||||||
@@ -48,43 +51,40 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
|
|||||||
|
|
||||||
hyperparameters = {
|
hyperparameters = {
|
||||||
"output_dir": [f"lora_models/{exp_name}"],
|
"output_dir": [f"lora_models/{exp_name}"],
|
||||||
"sd_model_version": ["sd15", "sdxl"],
|
"sd_model_version": ["sd15"],
|
||||||
"lora_training_urls": [
|
"lora_training_urls": [
|
||||||
"/home/rednax/Documents/datasets/grimes"
|
"/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean"
|
||||||
|
|
||||||
],
|
],
|
||||||
"concept_mode": ['face'],
|
"concept_mode": ['style'],
|
||||||
|
"sample_imgs_lora_scale": [0.9],
|
||||||
|
"caption_dropout": [0.2],
|
||||||
"seed": [0],
|
"seed": [0],
|
||||||
"resolution": [512],
|
"resolution": [512,640,768],
|
||||||
"train_batch_size": [4],
|
"train_batch_size": [8],
|
||||||
"n_sample_imgs": [6],
|
"n_sample_imgs": [8],
|
||||||
"max_train_steps": [400,800],
|
"max_train_steps": [2000],
|
||||||
"checkpointing_steps": [100],
|
"checkpointing_steps": [500],
|
||||||
"gradient_accumulation_steps": [1],
|
"gradient_accumulation_steps": [1],
|
||||||
|
|
||||||
"n_tokens": [2],
|
"n_tokens": [3],
|
||||||
"ti_lr": [0.001,0.0005],
|
"disable_ti": ['true'],
|
||||||
"ti_weight_decay": [0.001,0.0],
|
"ti_lr": [0.0001],
|
||||||
"l1_penalty": [0.0],
|
"token_warmup_steps": [0],
|
||||||
"token_warmup_steps": [0,60],
|
|
||||||
"tok_cov_reg_w": [2000],
|
"unet_lr": [0.001, 0.0003],
|
||||||
"cond_reg_w": [0.01e-5],
|
"lora_rank": [8,24,64],
|
||||||
"tok_cond_reg_w": [0.01e-5],
|
|
||||||
|
|
||||||
"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'],
|
"use_dora": ['false', 'true'],
|
||||||
|
|
||||||
|
"unet_optimizer_type": ['adamw'],
|
||||||
|
"is_lora": ['true'],
|
||||||
|
|
||||||
"text_encoder_lora_optimizer": [None],
|
"text_encoder_lora_optimizer": [None],
|
||||||
"text_encoder_lora_lr": [0.0e-4],
|
"text_encoder_lora_lr": [0.0e-4],
|
||||||
|
|
||||||
"snr_gamma": [5.0],
|
"snr_gamma": [5.0],
|
||||||
"caption_model": ["blip", "gpt4-v"],
|
"caption_model": ["florence", "blip", "no_caption"],
|
||||||
"augment_imgs_up_to_n": [20,40],
|
"augment_imgs_up_to_n": [40],
|
||||||
"verbose": ['true'],
|
"verbose": ['true'],
|
||||||
"debug": ['true']
|
"debug": ['true']
|
||||||
}
|
}
|
||||||
@@ -100,7 +100,7 @@ shutil.rmtree(config_output_dir, ignore_errors=True)
|
|||||||
os.makedirs(config_output_dir, exist_ok=True)
|
os.makedirs(config_output_dir, exist_ok=True)
|
||||||
|
|
||||||
# Open the shell script file
|
# Open the shell script file
|
||||||
try_sampling_n_times = 200
|
try_sampling_n_times = 120
|
||||||
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
|
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
|
||||||
resamples, combination = 0, None
|
resamples, combination = 0, None
|
||||||
|
|
||||||
@@ -122,9 +122,6 @@ for exp_index in tqdm(range(n_exp)): # number of combinations you want to gener
|
|||||||
dirname = os.path.dirname(config_filename)
|
dirname = os.path.dirname(config_filename)
|
||||||
os.makedirs(dirname, exist_ok=True)
|
os.makedirs(dirname, exist_ok=True)
|
||||||
|
|
||||||
# Make some final adjustments to the experiment settings before saving to disk:
|
|
||||||
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
|
|
||||||
|
|
||||||
with open(config_filename, "w") as f:
|
with open(config_filename, "w") as f:
|
||||||
json.dump(experiment_settings, f, indent=4)
|
json.dump(experiment_settings, f, indent=4)
|
||||||
break
|
break
|
||||||
@@ -146,7 +143,12 @@ def generate_sh_script(folder_path, output_sh_path):
|
|||||||
|
|
||||||
# Write a command for each JSON file
|
# Write a command for each JSON file
|
||||||
for json_file in json_files:
|
for json_file in json_files:
|
||||||
command = f"python main.py {os.path.join(folder_path, json_file)}\n"
|
file_path = os.path.join("scripts/", folder_path, json_file)
|
||||||
|
command = f"python main.py {file_path}\n"
|
||||||
|
|
||||||
|
if nohup:
|
||||||
|
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
|
||||||
|
|
||||||
sh_file.write(command)
|
sh_file.write(command)
|
||||||
|
|
||||||
generate_sh_script(config_output_dir, output_sh_path)
|
generate_sh_script(config_output_dir, output_sh_path)
|
||||||
|
|||||||
@@ -0,0 +1,199 @@
|
|||||||
|
import os
|
||||||
|
import json
|
||||||
|
import numpy as np
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
import seaborn as sns
|
||||||
|
from collections import defaultdict
|
||||||
|
from sklearn.metrics import r2_score
|
||||||
|
|
||||||
|
# Step 2: Define a function to count JPG files in a directory
|
||||||
|
def count_jpg_files(directory):
|
||||||
|
return len([f for f in os.listdir(directory) if f.lower().endswith('.jpg')])
|
||||||
|
|
||||||
|
# Step 3: Define a function to find and load the training_args.json file
|
||||||
|
def load_training_args(directory):
|
||||||
|
for root, dirs, files in os.walk(directory):
|
||||||
|
if 'training_args.json' in files:
|
||||||
|
with open(os.path.join(root, 'training_args.json'), 'r') as f:
|
||||||
|
return json.load(f)
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Step 4: Traverse the directory structure and collect data
|
||||||
|
def collect_data(root_dir, mode):
|
||||||
|
data = []
|
||||||
|
for root, dirs, files in os.walk(root_dir):
|
||||||
|
if "checkpoints" in dirs:
|
||||||
|
checkpoints_dir = os.path.join(root,"checkpoints")
|
||||||
|
|
||||||
|
if mode == "final_checkpoint":
|
||||||
|
checkpoint_dirs = sorted([f for f in os.listdir(checkpoints_dir) if os.path.isdir(os.path.join(checkpoints_dir, f))])
|
||||||
|
checkpoint_dir = os.path.join(checkpoints_dir, checkpoint_dirs[-1])
|
||||||
|
score = count_jpg_files(checkpoint_dir)
|
||||||
|
training_args = load_training_args(checkpoint_dir)
|
||||||
|
if mode == "n_validation_grids":
|
||||||
|
score = count_jpg_files(checkpoints_dir)
|
||||||
|
training_args = load_training_args(checkpoints_dir)
|
||||||
|
|
||||||
|
if not training_args:
|
||||||
|
print(f"Warning: Could not load training_args.json from {checkpoints_dir}")
|
||||||
|
continue
|
||||||
|
data.append((training_args, score))
|
||||||
|
|
||||||
|
print(f"Collected data from {len(data)} runs")
|
||||||
|
return data
|
||||||
|
|
||||||
|
# Step 5: Process the collected data to identify varying hyperparameters
|
||||||
|
def identify_varying_hyperparams(data, skip_params=['output_dir', 'start_time', 'name']):
|
||||||
|
all_params = set().union(*[set(args.keys()) for args, _ in data])
|
||||||
|
varying_params = {}
|
||||||
|
|
||||||
|
def make_hashable(val):
|
||||||
|
if isinstance(val, dict):
|
||||||
|
return tuple(sorted((k, make_hashable(v)) for k, v in val.items()))
|
||||||
|
elif isinstance(val, list):
|
||||||
|
return tuple(make_hashable(v) for v in val)
|
||||||
|
elif isinstance(val, set):
|
||||||
|
return frozenset(make_hashable(v) for v in val)
|
||||||
|
return val
|
||||||
|
|
||||||
|
for param in all_params:
|
||||||
|
if param in skip_params:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
values = [make_hashable(args.get(param)) for args, _ in data if param in args]
|
||||||
|
unique_values = set(values)
|
||||||
|
|
||||||
|
if len(unique_values) > 1:
|
||||||
|
# Check if all values are numeric
|
||||||
|
try:
|
||||||
|
numeric_values = [float(v) for v in unique_values]
|
||||||
|
varying_params[param] = set(numeric_values)
|
||||||
|
except ValueError:
|
||||||
|
# If not all numeric, keep as is
|
||||||
|
varying_params[param] = unique_values
|
||||||
|
|
||||||
|
print(f"---> Parameter '{param}' varies across runs")
|
||||||
|
|
||||||
|
# Special handling for dictionary-type parameters
|
||||||
|
if all(isinstance(v, dict) for v in values):
|
||||||
|
print(f"Dictionary values for '{param}':")
|
||||||
|
for v in unique_values:
|
||||||
|
print(f" {v}")
|
||||||
|
elif len(unique_values) <= 5: # Print up to 5 unique values
|
||||||
|
print(f"Unique values: {unique_values}")
|
||||||
|
else:
|
||||||
|
print(f"Number of unique values: {len(unique_values)}")
|
||||||
|
except TypeError as e:
|
||||||
|
print(f"Warning: Could not process values for parameter '{param}'. Error: {e}")
|
||||||
|
#print(f"Values: {[args.get(param) for args, _ in data if param in args]}")
|
||||||
|
|
||||||
|
return varying_params
|
||||||
|
|
||||||
|
def create_plots(data, varying_params, outdir, top = 0.15):
|
||||||
|
os.makedirs(outdir, exist_ok=True)
|
||||||
|
|
||||||
|
for param, values in varying_params.items():
|
||||||
|
if all(isinstance(v, dict) for v in values):
|
||||||
|
print(f"Skipping plot for dictionary parameter '{param}'")
|
||||||
|
continue
|
||||||
|
|
||||||
|
plt.figure(figsize=(12, 8))
|
||||||
|
param_data = defaultdict(list)
|
||||||
|
all_scores = []
|
||||||
|
|
||||||
|
for args, score in data:
|
||||||
|
if param in args:
|
||||||
|
value = args[param]
|
||||||
|
value_str = str(value)
|
||||||
|
param_data[value_str].append(score)
|
||||||
|
all_scores.append(score)
|
||||||
|
|
||||||
|
# Calculate global top threshold
|
||||||
|
global_top_percent = np.percentile(all_scores, 100*(1-top))
|
||||||
|
top_75_percent = np.percentile(all_scores, 25)
|
||||||
|
|
||||||
|
# Sort the values
|
||||||
|
try:
|
||||||
|
values_list = sorted(param_data.keys(), key=float)
|
||||||
|
except ValueError:
|
||||||
|
values_list = sorted(param_data.keys())
|
||||||
|
|
||||||
|
all_x = []
|
||||||
|
all_y = []
|
||||||
|
|
||||||
|
for i, value_str in enumerate(values_list):
|
||||||
|
scores = param_data[value_str]
|
||||||
|
|
||||||
|
# Apply jitter
|
||||||
|
jittered_x = np.random.normal(i, 0.1, size=len(scores))
|
||||||
|
jittered_y = np.array(scores) + np.random.normal(0, 0.01 * max(scores), size=len(scores))
|
||||||
|
|
||||||
|
# Use global top 25% threshold
|
||||||
|
top_mask = np.array(scores) >= global_top_percent
|
||||||
|
|
||||||
|
# Emphasize top 25% scores using JITTERED coordinates for plotting
|
||||||
|
sns.scatterplot(x=jittered_x[top_mask], y=jittered_y[top_mask], alpha=0.6, color='black', marker='X', s=40, linewidth=1)
|
||||||
|
|
||||||
|
# Plot all scores using JITTERED coordinates
|
||||||
|
sns.scatterplot(x=jittered_x, y=jittered_y, alpha=0.6, label=value_str)
|
||||||
|
|
||||||
|
all_x.extend([i] * len(scores))
|
||||||
|
all_y.extend(scores)
|
||||||
|
|
||||||
|
# Calculate trendline for top 75% data
|
||||||
|
x = np.array(all_x)
|
||||||
|
y = np.array(all_y)
|
||||||
|
top_75_mask = y >= top_75_percent
|
||||||
|
x_75 = x[top_75_mask]
|
||||||
|
y_75 = y[top_75_mask]
|
||||||
|
|
||||||
|
z = np.polyfit(x_75, y_75, 1)
|
||||||
|
p = np.poly1d(z)
|
||||||
|
plt.plot(range(len(values_list)), p(range(len(values_list))), "r--", alpha=0.8,
|
||||||
|
label=f'Top 75%: y={z[0]:.2f}x+{z[1]:.2f}\nR²: {r2_score(y_75, p(x_75)):.4f}')
|
||||||
|
|
||||||
|
# Calculate trendline for global top 25% scoring datapoints
|
||||||
|
top_mask = y >= global_top_percent
|
||||||
|
x_top = x[top_mask]
|
||||||
|
y_top = y[top_mask]
|
||||||
|
|
||||||
|
z_top = np.polyfit(x_top, y_top, 1)
|
||||||
|
p_top = np.poly1d(z_top)
|
||||||
|
plt.plot(range(len(values_list)), p_top(range(len(values_list))), "g--", alpha=0.8,
|
||||||
|
label=f'Top 25%: y={z_top[0]:.2f}x+{z_top[1]:.2f}\nR²: {r2_score(y_top, p_top(x_top)):.4f}')
|
||||||
|
|
||||||
|
plt.xlabel(param)
|
||||||
|
plt.ylabel('Score')
|
||||||
|
plt.title(f'Effect of {param} on Score')
|
||||||
|
|
||||||
|
# Set x-ticks and labels
|
||||||
|
plt.xticks(range(len(values_list)), values_list, rotation=45, ha='right')
|
||||||
|
|
||||||
|
# Adjust legend
|
||||||
|
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
|
||||||
|
|
||||||
|
plt.tight_layout()
|
||||||
|
|
||||||
|
# Save figure with error handling
|
||||||
|
try:
|
||||||
|
plt.savefig(f'{outdir}/{param}_vs_score.png', dpi=200, bbox_inches='tight')
|
||||||
|
except ValueError:
|
||||||
|
print(f"Warning: Failed to save image for {param}. Skipping...")
|
||||||
|
|
||||||
|
plt.close()
|
||||||
|
|
||||||
|
print(f"Plots have been saved as PNG files in {outdir}")
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
root_dir = "/home/rednax/SSD2TB/Github_repos/Eden/sd-lora-trainer/lora_models/ygor_sd15"
|
||||||
|
outdir = os.path.join('.', os.path.basename(root_dir))
|
||||||
|
|
||||||
|
# Collect data
|
||||||
|
#data = collect_data(root_dir, mode = "final_checkpoint")
|
||||||
|
data = collect_data(root_dir, mode = "n_validation_grids")
|
||||||
|
|
||||||
|
# Identify varying hyperparameters
|
||||||
|
varying_params = identify_varying_hyperparams(data)
|
||||||
|
|
||||||
|
# Create plots
|
||||||
|
create_plots(data, varying_params, outdir)
|
||||||
@@ -1,59 +0,0 @@
|
|||||||
"""
|
|
||||||
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!")
|
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
{
|
||||||
|
"name": "stitchly",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/stitchly_florence",
|
||||||
|
"concept_mode": "style",
|
||||||
|
"sample_imgs_lora_scale": 0.9,
|
||||||
|
"seed": 1,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 5000,
|
||||||
|
"checkpointing_steps": 1000,
|
||||||
|
|
||||||
|
"unet_optimizer_type": "AdamW8bit",
|
||||||
|
|
||||||
|
"disable_ti": true,
|
||||||
|
"is_lora": false,
|
||||||
|
"caption_dropout": 0.2,
|
||||||
|
"caption_model": "florence",
|
||||||
|
"ti_lr": 0.0001,
|
||||||
|
"unet_lr": 0.0001,
|
||||||
|
"lora_rank": 4,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
{
|
||||||
|
"unet_lr": 0.00,
|
||||||
|
"lora_rank": 4,
|
||||||
|
"freeze_ti_after_completion_f": 1.0,
|
||||||
|
|
||||||
|
"name": "mira_sdxl_ti",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
|
||||||
|
"concept_mode": "face",
|
||||||
|
"sample_imgs_lora_scale": 0.7,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 900,
|
||||||
|
"checkpointing_steps": 300,
|
||||||
|
|
||||||
|
"caption_model": "blip",
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "banny",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip",
|
||||||
|
"concept_mode": "object",
|
||||||
|
"sample_imgs_lora_scale": 0.75,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 300,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"disable_ti": false,
|
||||||
|
"caption_model": "florence",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0003,
|
||||||
|
"lora_rank": 16,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "xander_sd15",
|
||||||
|
"sd_model_version": "sd15",
|
||||||
|
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
|
||||||
|
"concept_mode": "face",
|
||||||
|
"sample_imgs_lora_scale": 0.8,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 768,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 600,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"disable_ti": false,
|
||||||
|
"caption_model": "florence",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0005,
|
||||||
|
"lora_rank": 16,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
{
|
||||||
|
"name": "xander_sdxl",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
|
||||||
|
"concept_mode": "face",
|
||||||
|
"sample_imgs_lora_scale": 0.7,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 300,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"caption_model": "blip",
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "banny",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip",
|
||||||
|
"concept_mode": "object",
|
||||||
|
"sample_imgs_lora_scale": 0.75,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 300,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"disable_ti": false,
|
||||||
|
"caption_model": "florence",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0003,
|
||||||
|
"lora_rank": 16,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "twisting_realities_sd15",
|
||||||
|
"sd_model_version": "sd15",
|
||||||
|
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean",
|
||||||
|
"concept_mode": "style",
|
||||||
|
"sample_imgs_lora_scale": 0.8,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 768,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 1200,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"disable_ti": false,
|
||||||
|
"caption_model": "blip",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0003,
|
||||||
|
"lora_rank": 16,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "ygor_painting_sd15",
|
||||||
|
"sd_model_version": "sd15",
|
||||||
|
"lora_training_urls": "/home/rednax/Documents/datasets/good_styles/visionary_painting_ygor_marotta_clean",
|
||||||
|
"concept_mode": "style",
|
||||||
|
"sample_imgs_lora_scale": 0.7,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 640,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 4000,
|
||||||
|
"checkpointing_steps": 500,
|
||||||
|
|
||||||
|
"disable_ti": true,
|
||||||
|
"caption_model": "blip",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0015,
|
||||||
|
"lora_rank": 64,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
{
|
||||||
|
"name": "twisting_realities",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
|
||||||
|
"concept_mode": "style",
|
||||||
|
"sample_imgs_lora_scale": 0.75,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 6,
|
||||||
|
"max_train_steps": 300,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"disable_ti": false,
|
||||||
|
"caption_model": "florence",
|
||||||
|
"ti_lr": 0.001,
|
||||||
|
"unet_lr": 0.0003,
|
||||||
|
"lora_rank": 16,
|
||||||
|
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
{
|
||||||
|
"name": "xander_sdxl",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "/home/rednax/Documents/datasets/beeple",
|
||||||
|
"concept_mode": "style",
|
||||||
|
"sample_imgs_lora_scale": 0.7,
|
||||||
|
"seed": 0,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 300,
|
||||||
|
"checkpointing_steps": 200,
|
||||||
|
|
||||||
|
"caption_model": "florence",
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
{
|
||||||
|
"name": "mira_sdxl",
|
||||||
|
"sd_model_version": "sdxl",
|
||||||
|
"lora_training_urls": "/home/rednax/Documents/datasets/01_people/mira",
|
||||||
|
"concept_mode": "face",
|
||||||
|
"sample_imgs_lora_scale": 0.7,
|
||||||
|
"seed": 1,
|
||||||
|
"resolution": 512,
|
||||||
|
"train_batch_size": 4,
|
||||||
|
"n_sample_imgs": 8,
|
||||||
|
"max_train_steps": 360,
|
||||||
|
"checkpointing_steps": 240,
|
||||||
|
|
||||||
|
"caption_model": "florence",
|
||||||
|
"debug": true
|
||||||
|
}
|
||||||
+37
-6
@@ -54,10 +54,31 @@ def set_adapter_scales(pipe, lora_scale = 1.0):
|
|||||||
|
|
||||||
return pipe
|
return pipe
|
||||||
|
|
||||||
|
import re
|
||||||
|
def remove_delimiter_characters(name: str, max_length: int = 255) -> str:
|
||||||
|
# Define a regular expression pattern to match all weird and special characters
|
||||||
|
pattern = r'[^\w.-]+' # Matches any character that is not alphanumeric, underscore, dot, or hyphen
|
||||||
|
|
||||||
|
# Replace all occurrences of the pattern with a single underscore
|
||||||
|
cleaned_name = re.sub(pattern, '_', name)
|
||||||
|
|
||||||
|
# Replace multiple consecutive underscores with a single underscore
|
||||||
|
cleaned_name = re.sub(r'_+', '_', cleaned_name)
|
||||||
|
|
||||||
|
# Strip leading or trailing underscores and dots
|
||||||
|
cleaned_name = cleaned_name.strip('_.')
|
||||||
|
|
||||||
|
# Ensure the name doesn't start with a dot (to avoid hidden files on Unix)
|
||||||
|
cleaned_name = cleaned_name.lstrip('.')
|
||||||
|
|
||||||
|
# Truncate to max_length if necessary
|
||||||
|
cleaned_name = cleaned_name[:max_length]
|
||||||
|
|
||||||
def remove_delimiter_characters(name: str):
|
# Raise an error if the name is empty or malformed after cleaning
|
||||||
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
|
if not cleaned_name:
|
||||||
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
|
raise ValueError("Malformed name")
|
||||||
|
|
||||||
|
return cleaned_name
|
||||||
|
|
||||||
# Convert to WebUI format
|
# Convert to WebUI format
|
||||||
def convert_pytorch_lora_safetensors_to_webui(
|
def convert_pytorch_lora_safetensors_to_webui(
|
||||||
@@ -135,7 +156,7 @@ def save_checkpoint(
|
|||||||
embedding_handler.save_embeddings(
|
embedding_handler.save_embeddings(
|
||||||
os.path.join(
|
os.path.join(
|
||||||
output_dir,
|
output_dir,
|
||||||
f"{name}_embeddings.safetensors"
|
f"{name}_{pretrained_model_version}_embeddings.safetensors"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -145,7 +166,7 @@ def save_checkpoint(
|
|||||||
output_dir, "special_params.json"
|
output_dir, "special_params.json"
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
if is_lora:
|
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"
|
assert len(unet_lora_parameters) > 0, f"Expected len(unet_lora_parameters) to be greater than zero if is_lora is True"
|
||||||
|
|
||||||
@@ -184,11 +205,21 @@ def save_checkpoint(
|
|||||||
|
|
||||||
convert_pytorch_lora_safetensors_to_webui(
|
convert_pytorch_lora_safetensors_to_webui(
|
||||||
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
|
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
|
||||||
output_filename=os.path.join(output_dir, f"{name}.safetensors")
|
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_lora.safetensors")
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
# Save the entire, finetuned unet weights:
|
||||||
unet.save_pretrained(save_directory = output_dir)
|
unet.save_pretrained(save_directory = output_dir)
|
||||||
|
|
||||||
|
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
|
||||||
|
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
|
||||||
|
for file in to_remove:
|
||||||
|
file_path = os.path.join(output_dir, file)
|
||||||
|
if os.path.exists(file_path):
|
||||||
|
os.remove(file_path)
|
||||||
|
|
||||||
|
return
|
||||||
|
|
||||||
def load_checkpoint(
|
def load_checkpoint(
|
||||||
pretrained_model_version: str,
|
pretrained_model_version: str,
|
||||||
pretrained_model_path: str,
|
pretrained_model_path: str,
|
||||||
|
|||||||
+67
-34
@@ -3,15 +3,47 @@ from datetime import datetime
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
import json, time, os
|
import json, time, os
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
from trainer.models import pretrained_models
|
|
||||||
from trainer.utils.utils import pick_best_gpu_id
|
from trainer.utils.utils import pick_best_gpu_id
|
||||||
|
from trainer.checkpoint import remove_delimiter_characters
|
||||||
|
|
||||||
|
class ModelPaths:
|
||||||
|
def __init__(self):
|
||||||
|
self.paths = {
|
||||||
|
"BLIP": "./cache",
|
||||||
|
"FLORENCE": "./cache",
|
||||||
|
"CLIP": "./cache",
|
||||||
|
"SR": "./cache",
|
||||||
|
"SD": "./models",
|
||||||
|
}
|
||||||
|
|
||||||
|
def get_path(self, key):
|
||||||
|
return self.paths.get(key, None)
|
||||||
|
|
||||||
|
def set_path(self, key, path):
|
||||||
|
if key in self.paths:
|
||||||
|
self.paths[key] = path
|
||||||
|
|
||||||
|
model_paths = ModelPaths()
|
||||||
|
|
||||||
|
# Default SD model download urls in case no local model is found:
|
||||||
|
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/Eden_SDXL.safetensors"
|
||||||
|
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-XL-v6/resolve/main/juggernautXL_version6Rundiffusion.safetensors"
|
||||||
|
SD15_URL = "https://huggingface.co/KamCastle/jugg/resolve/main/juggernaut_reborn.safetensors"
|
||||||
|
|
||||||
|
pretrained_models = {
|
||||||
|
"sdxl": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SDXL_URL)), "url": SDXL_URL, "version": "sdxl"},
|
||||||
|
"sd15": {"path": os.path.join(model_paths.get_path("SD"), os.path.basename(SD15_URL)), "url": SD15_URL, "version": "sd15"}
|
||||||
|
}
|
||||||
|
|
||||||
class TrainingConfig(BaseModel):
|
class TrainingConfig(BaseModel):
|
||||||
lora_training_urls: str
|
lora_training_urls: str
|
||||||
concept_mode: Literal["face", "style", "object"]
|
concept_mode: Literal["face", "style", "object"]
|
||||||
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
|
caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing
|
||||||
caption_model: Literal["gpt4-v", "blip"] = "blip"
|
prompt_modifier: str = None # optional prompt modifier
|
||||||
sd_model_version: Literal["sdxl", "sd15", "sd3"]
|
caption_model: Literal["gpt4-v", "blip", "florence", "no_caption"] = "florence"
|
||||||
|
caption_dropout: float = 0.1 # dropout rate for captions: occasionally use empty prompt
|
||||||
|
sd_model_version: Literal["sdxl", "sd15", None] = None
|
||||||
|
ckpt_path: str = None # optional hardcoded checkpoint path
|
||||||
pretrained_model: dict = None
|
pretrained_model: dict = None
|
||||||
seed: Union[int, None] = None
|
seed: Union[int, None] = None
|
||||||
resolution: int = 512
|
resolution: int = 512
|
||||||
@@ -19,36 +51,36 @@ class TrainingConfig(BaseModel):
|
|||||||
train_img_size: List[int] = None
|
train_img_size: List[int] = None
|
||||||
train_aspect_ratio: float = None
|
train_aspect_ratio: float = None
|
||||||
train_batch_size: int = 4
|
train_batch_size: int = 4
|
||||||
num_train_epochs: int = 10000
|
max_train_steps: int = 300
|
||||||
max_train_steps: int = 360
|
num_train_epochs: int = None
|
||||||
checkpointing_steps: int = 10000
|
checkpointing_steps: int = 10000
|
||||||
gradient_accumulation_steps: int = 1
|
gradient_accumulation_steps: int = 1
|
||||||
is_lora: bool = True
|
is_lora: bool = True
|
||||||
|
|
||||||
unet_optimizer_type: Literal["adamw", "prodigy", "adamw_8bit"] = "adamw"
|
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
|
||||||
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
|
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
|
||||||
unet_lr: float = 1.0e-3
|
unet_lr: float = 0.0003
|
||||||
prodigy_d_coef: float = 1.0
|
prodigy_d_coef: float = 1.0
|
||||||
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
|
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
|
||||||
lora_weight_decay: float = 0.002
|
lora_weight_decay: float = 0.004
|
||||||
# if ti_lr is None, then we completely skip textual inversion
|
|
||||||
ti_lr: Union[float, None] = 1e-3
|
ti_lr: float = 0.001
|
||||||
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
|
|
||||||
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
|
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
|
||||||
ti_weight_decay: float = 0.0
|
ti_weight_decay: float = 0.0
|
||||||
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
|
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
|
||||||
freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done
|
freeze_ti_after_completion_f: float = 0.7 # freeze the TI after this fraction of the training is done
|
||||||
|
freeze_unet_before_completion_f: float = 0.0 # freeze the UNET before this fraction of the training is done
|
||||||
|
|
||||||
|
token_attention_loss_w: float = 3e-7
|
||||||
cond_reg_w: float = 0.0e-5
|
cond_reg_w: float = 0.0e-5
|
||||||
tok_cond_reg_w: float = 0.0e-5
|
tok_cond_reg_w: float = 0.0e-5
|
||||||
tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens
|
tok_cov_reg_w: float = 0. # regularizes the token covariance matrix wrt pretrained, normal tokens
|
||||||
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std
|
l1_penalty: float = 0.03 # Makes the unet lora matrix more sparse
|
||||||
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
|
|
||||||
|
|
||||||
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
|
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
|
||||||
snr_gamma: float = 5.0
|
snr_gamma: float = 5.0
|
||||||
lora_alpha_multiplier: float = 1.0
|
lora_alpha_multiplier: float = 1.0
|
||||||
lora_rank: int = 12
|
lora_rank: int = 16
|
||||||
use_dora: bool = False
|
use_dora: bool = False
|
||||||
|
|
||||||
left_right_flip_augmentation: bool = True
|
left_right_flip_augmentation: bool = True
|
||||||
@@ -59,22 +91,17 @@ class TrainingConfig(BaseModel):
|
|||||||
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
|
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
|
||||||
n_sample_imgs: int = 4
|
n_sample_imgs: int = 4
|
||||||
name: str = None
|
name: str = None
|
||||||
output_dir: str = "lora_models/unnamed"
|
output_dir: str = "eden_lora_training_runs"
|
||||||
debug: bool = False
|
debug: bool = False
|
||||||
allow_tf32: bool = True
|
allow_tf32: bool = True
|
||||||
remove_ti_token_from_prompts: bool = False
|
disable_ti: bool = False
|
||||||
|
skip_gpt_cleanup: bool = False
|
||||||
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
|
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
|
||||||
n_tokens: int = 2
|
n_tokens: int = 3
|
||||||
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
|
inserting_list_tokens: List[str] = ["<s0>","<s1>","<s2>"]
|
||||||
token_dict: dict = {"TOK": "<s0><s1>"}
|
token_dict: dict = {"TOK": "<s0><s1><s2>"}
|
||||||
device: str = "cuda:0"
|
device: str = "cuda:0"
|
||||||
crops_coords_top_left_h: int = 0
|
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
|
||||||
crops_coords_top_left_w: int = 0
|
|
||||||
do_cache: bool = True
|
|
||||||
unet_learning_rate: float = 1.0
|
|
||||||
lr_num_cycles: int = 1
|
|
||||||
lr_power: float = 1.0
|
|
||||||
sample_imgs_lora_scale: float = 0.65 # Default lora scale for sampling the validation images
|
|
||||||
dataloader_num_workers: int = 0
|
dataloader_num_workers: int = 0
|
||||||
training_attributes: dict = {}
|
training_attributes: dict = {}
|
||||||
aspect_ratio_bucketing: bool = False
|
aspect_ratio_bucketing: bool = False
|
||||||
@@ -93,16 +120,19 @@ class TrainingConfig(BaseModel):
|
|||||||
|
|
||||||
def __init__(self, **data):
|
def __init__(self, **data):
|
||||||
super().__init__(**data)
|
super().__init__(**data)
|
||||||
self.pretrained_model = pretrained_models[self.sd_model_version]
|
|
||||||
|
|
||||||
# add some metrics to the foldername:
|
if not self.ckpt_path:
|
||||||
lora_str = "dora" if self.use_dora else "lora"
|
self.pretrained_model = pretrained_models[self.sd_model_version]
|
||||||
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
|
else:
|
||||||
|
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
|
||||||
|
|
||||||
if not self.name:
|
if not self.name:
|
||||||
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
|
self.name = os.path.basename(self.lora_training_urls)[:40]
|
||||||
|
|
||||||
self.output_dir = self.output_dir + f"--{timestamp_short}-{self.sd_model_version}_{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
|
self.name = remove_delimiter_characters(self.name)
|
||||||
|
|
||||||
|
timestamp = datetime.now().strftime("%d%b_%H%M")
|
||||||
|
self.output_dir = self.output_dir + f"/{self.name}_{timestamp}-{self.concept_mode}_res{self.resolution}_{self.max_train_steps}steps"
|
||||||
os.makedirs(self.output_dir, exist_ok=True)
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
|
|
||||||
if self.seed is None:
|
if self.seed is None:
|
||||||
@@ -111,6 +141,9 @@ class TrainingConfig(BaseModel):
|
|||||||
if self.unet_lr_warmup_steps is None:
|
if self.unet_lr_warmup_steps is None:
|
||||||
self.unet_lr_warmup_steps = self.max_train_steps
|
self.unet_lr_warmup_steps = self.max_train_steps
|
||||||
|
|
||||||
|
if self.checkpointing_steps < 1:
|
||||||
|
self.checkpointing_steps = self.max_train_steps
|
||||||
|
|
||||||
if self.concept_mode == "face":
|
if self.concept_mode == "face":
|
||||||
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
|
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
|
||||||
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
|
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
|
||||||
|
|||||||
+34
-31
@@ -1,12 +1,12 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
from tqdm import tqdm
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import PIL
|
import PIL
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
from typing import Tuple, Dict, List
|
from typing import Tuple, Dict, List
|
||||||
from tqdm import tqdm
|
|
||||||
|
|
||||||
def prepare_image(
|
def prepare_image(
|
||||||
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
|
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
|
||||||
@@ -33,7 +33,6 @@ class PreprocessedDataset(Dataset):
|
|||||||
data_dir: str,
|
data_dir: str,
|
||||||
pipe,
|
pipe,
|
||||||
vae_encoder,
|
vae_encoder,
|
||||||
do_cache: bool = False,
|
|
||||||
size: List[int] = [512, 512],
|
size: List[int] = [512, 512],
|
||||||
text_dropout: float = 0.0,
|
text_dropout: float = 0.0,
|
||||||
aspect_ratio_bucketing: bool = False,
|
aspect_ratio_bucketing: bool = False,
|
||||||
@@ -43,13 +42,14 @@ class PreprocessedDataset(Dataset):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.data_dir = data_dir
|
self.data_dir = data_dir
|
||||||
self.csv_path = os.path.join(data_dir, "captions.csv")
|
self.csv_path = os.path.join(data_dir, "captions.csv")
|
||||||
self.data = pd.read_csv(self.csv_path)
|
self.data = pd.read_csv(self.csv_path, dtype={"caption": str})
|
||||||
|
|
||||||
self.captions = self.data["caption"]
|
self.captions = self.data["caption"]
|
||||||
self.captions = self.captions.str.lower()
|
self.captions = self.captions.str.lower()
|
||||||
for key, value in substitute_caption_map.items():
|
for key, value in substitute_caption_map.items():
|
||||||
self.captions = self.captions.str.replace(key.lower(), value)
|
self.captions = self.captions.str.replace(key.lower(), value)
|
||||||
|
|
||||||
|
self.captions = self.captions.fillna("")
|
||||||
self.image_path = self.data["image_path"]
|
self.image_path = self.data["image_path"]
|
||||||
|
|
||||||
if "mask_path" not in self.data.columns:
|
if "mask_path" not in self.data.columns:
|
||||||
@@ -63,23 +63,31 @@ class PreprocessedDataset(Dataset):
|
|||||||
self.text_dropout = text_dropout
|
self.text_dropout = text_dropout
|
||||||
self.size = size
|
self.size = size
|
||||||
|
|
||||||
if do_cache:
|
# If the training data is small we can keep everything in memory, otherwise offload to disk
|
||||||
print("Caching latents, masks and captions...\n")
|
self.do_cache = True if len(self.data) < 500 else False
|
||||||
|
|
||||||
|
if self.do_cache:
|
||||||
|
print("Encoding latents, masks and captions and storing in memory...\n")
|
||||||
self.vae_latents = []
|
self.vae_latents = []
|
||||||
self.masks = []
|
self.masks = []
|
||||||
self.do_cache = True
|
|
||||||
|
|
||||||
for idx in tqdm(range(len(self.data))):
|
for idx in tqdm(range(len(self.data))):
|
||||||
if len(self.data) < 25:
|
vae_latent, mask, _ = self._process(idx)
|
||||||
print(self.captions[idx])
|
|
||||||
vae_latent, mask = self._process(idx)
|
|
||||||
self.vae_latents.append(vae_latent)
|
self.vae_latents.append(vae_latent)
|
||||||
self.masks.append(mask)
|
self.masks.append(mask.detach())
|
||||||
|
|
||||||
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
|
else: # Store the latents and masks on disk
|
||||||
del self.vae_encoder
|
print("Encoding latents, masks and captions and storing on disk...\n")
|
||||||
else:
|
self.vae_latents = None
|
||||||
self.do_cache = False
|
self.masks = None
|
||||||
|
|
||||||
|
for idx in tqdm(range(len(self.data))):
|
||||||
|
vae_latent, mask, image_path = self._process(idx)
|
||||||
|
torch.save(vae_latent, os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
|
||||||
|
torch.save(mask, os.path.join(self.data_dir, f"{idx}_mask.pt"))
|
||||||
|
|
||||||
|
del self.vae_encoder
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
if aspect_ratio_bucketing:
|
if aspect_ratio_bucketing:
|
||||||
print("Using aspect ratio bucketing.")
|
print("Using aspect ratio bucketing.")
|
||||||
@@ -101,12 +109,9 @@ class PreprocessedDataset(Dataset):
|
|||||||
def get_aspect_ratio_bucketed_batch(self):
|
def get_aspect_ratio_bucketed_batch(self):
|
||||||
assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()"
|
assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()"
|
||||||
indices, resolution = self.bucket_manager.get_batch()
|
indices, resolution = self.bucket_manager.get_batch()
|
||||||
|
|
||||||
print(f"Got bucket batch: {indices}, resolution: {resolution}")
|
|
||||||
tok1, tok2, vae_latents, masks = [], [], [], []
|
tok1, tok2, vae_latents, masks = [], [], [], []
|
||||||
|
|
||||||
for idx in indices:
|
for idx in indices:
|
||||||
|
|
||||||
if self.tokenizer_2 is None:
|
if self.tokenizer_2 is None:
|
||||||
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
|
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
|
||||||
else:
|
else:
|
||||||
@@ -142,28 +147,24 @@ class PreprocessedDataset(Dataset):
|
|||||||
image = PIL.Image.open(image_path).convert("RGB")
|
image = PIL.Image.open(image_path).convert("RGB")
|
||||||
if bucketing_resolution is None:
|
if bucketing_resolution is None:
|
||||||
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to(
|
image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to(
|
||||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
dtype=self.vae_encoder.dtype
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
|
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
|
||||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
dtype=self.vae_encoder.dtype
|
||||||
)
|
)
|
||||||
|
|
||||||
vae_latent = self.vae_encoder.encode(image).latent_dist
|
vae_latent = self.vae_encoder.encode(image.to(self.vae_encoder.device)).latent_dist
|
||||||
dummy_vae_latent = vae_latent.sample()
|
dummy_vae_latent = vae_latent.sample()
|
||||||
|
|
||||||
if self.mask_path is None:
|
if self.mask_path is None:
|
||||||
mask = torch.ones_like(
|
mask = torch.ones_like(dummy_vae_latent, dtype=self.vae_encoder.dtype)
|
||||||
dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
|
||||||
)
|
|
||||||
|
|
||||||
else:
|
else:
|
||||||
mask_path = self.mask_path[idx]
|
mask_path = self.mask_path[idx]
|
||||||
mask_path = os.path.join(self.data_dir, mask_path)
|
mask_path = os.path.join(self.data_dir, mask_path)
|
||||||
mask = PIL.Image.open(mask_path)
|
mask = PIL.Image.open(mask_path)
|
||||||
mask = prepare_mask(mask, self.size[0], self.size[1]).to(
|
mask = prepare_mask(mask, self.size[0], self.size[1]).to(dtype=self.vae_encoder.dtype)
|
||||||
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
|
|
||||||
)
|
|
||||||
|
|
||||||
mask_dtype = mask.dtype
|
mask_dtype = mask.dtype
|
||||||
mask = mask.float()
|
mask = mask.float()
|
||||||
@@ -175,7 +176,7 @@ class PreprocessedDataset(Dataset):
|
|||||||
|
|
||||||
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
|
assert len(mask.shape) == 4 and len(dummy_vae_latent.shape) == 4
|
||||||
|
|
||||||
return vae_latent, mask.squeeze()
|
return vae_latent, mask.squeeze(), image_path
|
||||||
|
|
||||||
def __getitem__(
|
def __getitem__(
|
||||||
self, idx: int, bucketing_resolution:tuple = None
|
self, idx: int, bucketing_resolution:tuple = None
|
||||||
@@ -183,10 +184,12 @@ class PreprocessedDataset(Dataset):
|
|||||||
|
|
||||||
if self.do_cache:
|
if self.do_cache:
|
||||||
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
|
vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor
|
||||||
return self.captions[idx], vae_latent.squeeze(), self.masks[idx]
|
return self.captions[idx], vae_latent.squeeze().detach(), self.masks[idx].detach()
|
||||||
else: # This code pathway has not been tested in a long time and might be broken
|
else: # Load from disk:
|
||||||
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution)
|
vae_latent = torch.load(os.path.join(self.data_dir, f"{idx}_vae_latent.pt"))
|
||||||
vae_latent = vae_latent.sample() * self.vae_scaling_factor
|
vae_latent = vae_latent.sample() * self.vae_scaling_factor
|
||||||
return caption, vae_latent.squeeze(), mask
|
mask = torch.load(os.path.join(self.data_dir, f"{idx}_mask.pt"))
|
||||||
|
caption = self.captions[idx]
|
||||||
|
return caption, vae_latent.squeeze().detach(), mask.detach()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+24
-123
@@ -9,7 +9,6 @@ from typing import List, Optional, Dict
|
|||||||
from safetensors.torch import save_file, safe_open
|
from safetensors.torch import save_file, safe_open
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
|
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
|
||||||
from transformers import T5EncoderModel
|
|
||||||
|
|
||||||
class TokenEmbeddingsHandler:
|
class TokenEmbeddingsHandler:
|
||||||
def __init__(self, text_encoders, tokenizers):
|
def __init__(self, text_encoders, tokenizers):
|
||||||
@@ -32,10 +31,7 @@ class TokenEmbeddingsHandler:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
# Directly accessing and modifying the original weights tensor
|
# Directly accessing and modifying the original weights tensor
|
||||||
if isinstance(text_encoder, T5EncoderModel):
|
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
|
||||||
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.")
|
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
|
||||||
|
|
||||||
def get_trainable_embeddings(self):
|
def get_trainable_embeddings(self):
|
||||||
@@ -53,22 +49,10 @@ class TokenEmbeddingsHandler:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
|
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
|
||||||
if isinstance(text_encoder, T5EncoderModel):
|
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
|
||||||
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
|
embeddings[f'txt_encoder_{idx}'] = token_embeddings
|
||||||
|
|
||||||
# Get all corresponding tokens for these embeddings
|
# Get all corresponding tokens for these embeddings
|
||||||
@@ -208,19 +192,9 @@ class TokenEmbeddingsHandler:
|
|||||||
self.non_train_ids = all_indices[inu]
|
self.non_train_ids = all_indices[inu]
|
||||||
|
|
||||||
# random initialization of new tokens
|
# random initialization of new tokens
|
||||||
"""
|
std_token_embedding = (
|
||||||
handle both T5EncoderModel and other text encoders
|
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
|
||||||
|
)
|
||||||
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
|
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
|
||||||
|
|
||||||
if starting_toks is not None:
|
if starting_toks is not None:
|
||||||
@@ -233,28 +207,14 @@ class TokenEmbeddingsHandler:
|
|||||||
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
|
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
|
||||||
else:
|
else:
|
||||||
std_multiplier = 1.0
|
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()
|
current_std = init_embeddings.std(dim=1).mean()
|
||||||
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
|
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()
|
||||||
|
|
||||||
if isinstance(text_encoder, T5EncoderModel):
|
self.embeddings_settings[
|
||||||
text_encoder.encoder.embed_tokens.weight.data[self.train_ids] = init_embeddings.clone()
|
f"original_embeddings_{idx}"
|
||||||
else:
|
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
|
||||||
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 = torch.ones((len(tokenizer),), dtype=torch.bool)
|
||||||
inu[self.train_ids] = False
|
inu[self.train_ids] = False
|
||||||
@@ -300,11 +260,7 @@ class TokenEmbeddingsHandler:
|
|||||||
# original_size = (config.resolution, config.resolution)
|
# original_size = (config.resolution, config.resolution)
|
||||||
original_size = (1024, 1024)
|
original_size = (1024, 1024)
|
||||||
target_size = (config.resolution, config.resolution)
|
target_size = (config.resolution, config.resolution)
|
||||||
|
crops_coords_top_left = (0,0)
|
||||||
crops_coords_top_left = (
|
|
||||||
config.crops_coords_top_left_h,
|
|
||||||
config.crops_coords_top_left_w,
|
|
||||||
)
|
|
||||||
|
|
||||||
if pipe.text_encoder_2 is None:
|
if pipe.text_encoder_2 is None:
|
||||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||||
@@ -437,7 +393,6 @@ class TokenEmbeddingsHandler:
|
|||||||
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
|
embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0.
|
||||||
|
|
||||||
optimizer_ti.step()
|
optimizer_ti.step()
|
||||||
self.fix_embedding_std(config.off_ratio_power)
|
|
||||||
optimizer_ti.zero_grad()
|
optimizer_ti.zero_grad()
|
||||||
|
|
||||||
if config.debug:
|
if config.debug:
|
||||||
@@ -454,27 +409,14 @@ class TokenEmbeddingsHandler:
|
|||||||
for idx, text_encoder in enumerate(self.text_encoders):
|
for idx, text_encoder in enumerate(self.text_encoders):
|
||||||
if text_encoder is None:
|
if text_encoder is None:
|
||||||
continue
|
continue
|
||||||
|
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
|
||||||
if isinstance(text_encoder, T5EncoderModel):
|
0
|
||||||
|
] == len(self.tokenizers[0]), "Tokenizers should be the same."
|
||||||
assert text_encoder.encoder.embed_tokens.weight.data.shape[
|
new_token_embeddings = (
|
||||||
0
|
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||||
] == len(self.tokenizers[idx]), "Tokenizers should be the same."
|
self.train_ids
|
||||||
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
|
tensors[txt_encoder_keys[idx]] = new_token_embeddings
|
||||||
|
|
||||||
save_file(tensors, file_path)
|
save_file(tensors, file_path)
|
||||||
@@ -483,41 +425,6 @@ class TokenEmbeddingsHandler:
|
|||||||
def device(self):
|
def device(self):
|
||||||
return self.text_encoders[0].device
|
return self.text_encoders[0].device
|
||||||
|
|
||||||
def fix_embedding_std(self, off_ratio_power=0.1):
|
|
||||||
if off_ratio_power == 0.0:
|
|
||||||
return
|
|
||||||
|
|
||||||
idx = 0
|
|
||||||
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
|
|
||||||
if text_encoder is None:
|
|
||||||
idx += 1
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Get the standard deviation target and current embeddings.
|
|
||||||
target_std = self.embeddings_settings[f"std_token_embedding_{idx}"]
|
|
||||||
embeddings, _ = self.get_trainable_embeddings()
|
|
||||||
new_embeddings = embeddings[f'txt_encoder_{idx}']
|
|
||||||
assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!"
|
|
||||||
|
|
||||||
new_stds = new_embeddings.std(dim=1)
|
|
||||||
#off_ratios = target_std.float() / new_stds.float()
|
|
||||||
off_ratios = target_std / new_stds
|
|
||||||
|
|
||||||
# Check if off_ratios are within an acceptable range.
|
|
||||||
if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1):
|
|
||||||
# Convert the pytorch tensor into a list of python floats:
|
|
||||||
off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3)
|
|
||||||
print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...")
|
|
||||||
|
|
||||||
# Adjust embeddings using the computed ratios.
|
|
||||||
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
|
|
||||||
index_updates = ~index_no_updates
|
|
||||||
multiplier_values = off_ratios**off_ratio_power
|
|
||||||
multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings)
|
|
||||||
text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values
|
|
||||||
|
|
||||||
idx += 1
|
|
||||||
|
|
||||||
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
|
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
|
||||||
# Assuming new tokens are of the format <s_i>
|
# Assuming new tokens are of the format <s_i>
|
||||||
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
|
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
|
||||||
@@ -527,15 +434,9 @@ class TokenEmbeddingsHandler:
|
|||||||
|
|
||||||
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
|
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."
|
assert self.train_ids is not None, "New tokens could not be converted to IDs."
|
||||||
|
text_encoder.text_model.embeddings.token_embedding.weight.data[
|
||||||
if isinstance(text_encoder, T5EncoderModel):
|
self.train_ids
|
||||||
text_encoder.encoder.embed_tokens.weight.data[
|
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
|
||||||
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"]):
|
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
|
||||||
if not os.path.exists(file_path):
|
if not os.path.exists(file_path):
|
||||||
|
|||||||
+10
-6
@@ -155,11 +155,7 @@ def get_conditioning_signals(config, pipe, captions):
|
|||||||
# original_size = (config.resolution, config.resolution)
|
# original_size = (config.resolution, config.resolution)
|
||||||
original_size = (1024, 1024)
|
original_size = (1024, 1024)
|
||||||
target_size = (config.resolution, config.resolution)
|
target_size = (config.resolution, config.resolution)
|
||||||
|
crops_coords_top_left = (0,0)
|
||||||
crops_coords_top_left = (
|
|
||||||
config.crops_coords_top_left_h,
|
|
||||||
config.crops_coords_top_left_w,
|
|
||||||
)
|
|
||||||
|
|
||||||
if pipe.text_encoder_2 is None:
|
if pipe.text_encoder_2 is None:
|
||||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||||
@@ -245,7 +241,7 @@ def encode_prompt_advanced(
|
|||||||
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
|
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
|
||||||
This allows interpolating the strength of the trained token in the final image.
|
This allows interpolating the strength of the trained token in the final image.
|
||||||
"""
|
"""
|
||||||
if lora_path:
|
if lora_path and token_scale != 0:
|
||||||
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
|
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
|
||||||
else:
|
else:
|
||||||
lora_prompt = prompt
|
lora_prompt = prompt
|
||||||
@@ -299,6 +295,8 @@ def render_images(
|
|||||||
is_lora,
|
is_lora,
|
||||||
pretrained_model,
|
pretrained_model,
|
||||||
lora_scale,
|
lora_scale,
|
||||||
|
disable_ti=False,
|
||||||
|
prompt_modifier=None,
|
||||||
n_steps=25,
|
n_steps=25,
|
||||||
n_imgs=4,
|
n_imgs=4,
|
||||||
device="cuda:0",
|
device="cuda:0",
|
||||||
@@ -328,6 +326,9 @@ def render_images(
|
|||||||
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
|
validation_prompts_raw = random.sample(val_prompts["object"], n_imgs)
|
||||||
validation_prompts_raw[0] = "<concept>"
|
validation_prompts_raw[0] = "<concept>"
|
||||||
|
|
||||||
|
if prompt_modifier:
|
||||||
|
validation_prompts_raw = [prompt_modifier.format(prompt) for prompt in validation_prompts_raw]
|
||||||
|
|
||||||
if (
|
if (
|
||||||
checkpoint_folder is not None
|
checkpoint_folder is not None
|
||||||
): # reload the entire pipeline from disk and load in the lora module
|
): # reload the entire pipeline from disk and load in the lora module
|
||||||
@@ -376,6 +377,7 @@ def render_images(
|
|||||||
lora_scale,
|
lora_scale,
|
||||||
guidance_scale=8,
|
guidance_scale=8,
|
||||||
concept_mode=concept_mode,
|
concept_mode=concept_mode,
|
||||||
|
token_scale = 0 if disable_ti else None
|
||||||
)
|
)
|
||||||
|
|
||||||
pipeline_args["prompt_embeds"] = c
|
pipeline_args["prompt_embeds"] = c
|
||||||
@@ -415,6 +417,7 @@ def render_images_eval(
|
|||||||
pretrained_model: dict,
|
pretrained_model: dict,
|
||||||
trigger_text: str,
|
trigger_text: str,
|
||||||
lora_scale=0.7,
|
lora_scale=0.7,
|
||||||
|
disable_ti=False,
|
||||||
n_steps=25,
|
n_steps=25,
|
||||||
n_imgs=4,
|
n_imgs=4,
|
||||||
device="cuda:0",
|
device="cuda:0",
|
||||||
@@ -469,6 +472,7 @@ def render_images_eval(
|
|||||||
lora_scale,
|
lora_scale,
|
||||||
guidance_scale=8,
|
guidance_scale=8,
|
||||||
concept_mode=concept_mode,
|
concept_mode=concept_mode,
|
||||||
|
token_scale = 0 if disable_ti else None
|
||||||
)
|
)
|
||||||
|
|
||||||
pipeline_args["prompt_embeds"] = c
|
pipeline_args["prompt_embeds"] = c
|
||||||
|
|||||||
+78
-15
@@ -4,7 +4,81 @@ import matplotlib.pyplot as plt
|
|||||||
import torch
|
import torch
|
||||||
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
|
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
|
||||||
from trainer.inference import get_conditioning_signals
|
from trainer.inference import get_conditioning_signals
|
||||||
from transformers import T5EncoderModel
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_loss, verbose=0):
|
||||||
|
"""
|
||||||
|
Custom loss function to regularize the attention maps of the token embeddings.
|
||||||
|
"""
|
||||||
|
masks = masks[:, 0].float()
|
||||||
|
img_ratio = masks.shape[-1] / masks.shape[-2]
|
||||||
|
|
||||||
|
att_L2_losses = []
|
||||||
|
ti_heatmaps = []
|
||||||
|
ti_masks = []
|
||||||
|
att_reg_threshold = 0.0
|
||||||
|
|
||||||
|
# attention_maps.shape = [n_layers, batch_size, w, h, 77]
|
||||||
|
attention_maps = daam_loss.process_and_stack_attention_scores(img_ratio)
|
||||||
|
n_layers, batch_size, w, h, n_tokens = attention_maps.shape
|
||||||
|
|
||||||
|
# reshape masks to match attention maps:
|
||||||
|
# masks.shape = [batch_size, w2, h2]
|
||||||
|
masks = F.interpolate(masks.unsqueeze(1), size=(attention_maps.shape[-3], attention_maps.shape[-2])).squeeze(1)
|
||||||
|
masks = masks.unsqueeze(0).unsqueeze(-1)
|
||||||
|
masks = masks.repeat(n_layers, 1, 1, 1, n_tokens)
|
||||||
|
|
||||||
|
for batch_index, caption in enumerate(captions):
|
||||||
|
token_indices = pipe.tokenizer.encode(caption)
|
||||||
|
|
||||||
|
# Penalize the mean attention score of each token:
|
||||||
|
mean_att_per_token = attention_maps[:,batch_index, :, :, 1:len(token_indices)-1].mean(dim=[0,1,2])
|
||||||
|
att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean()
|
||||||
|
att_L2_losses.append(att_L2_loss)
|
||||||
|
|
||||||
|
try:
|
||||||
|
ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids]
|
||||||
|
except:
|
||||||
|
continue
|
||||||
|
batch_ti_heatmaps, batch_ti_masks = [], []
|
||||||
|
# Extract the attention heatmaps corresponding to the trainable token embeddings:
|
||||||
|
for text_token_index in ti_token_indices:
|
||||||
|
ti_heatmap = attention_maps[:,batch_index, :, :, text_token_index].mean(dim=0)
|
||||||
|
ti_mask = masks[:,batch_index, :, :, text_token_index].mean(dim=0)
|
||||||
|
batch_ti_heatmaps.append(ti_heatmap.float())
|
||||||
|
batch_ti_masks.append(ti_mask)
|
||||||
|
|
||||||
|
ti_heatmaps.append(torch.stack(batch_ti_heatmaps))
|
||||||
|
ti_masks.append(torch.stack(batch_ti_masks))
|
||||||
|
|
||||||
|
if len(ti_heatmaps) == 0:
|
||||||
|
return torch.tensor(0.0).to(masks.dtype)
|
||||||
|
|
||||||
|
ti_heatmaps = torch.stack(ti_heatmaps)
|
||||||
|
ti_masks = torch.stack(ti_masks)
|
||||||
|
|
||||||
|
#ti_heatmaps.shape = [batch_size, n_tokens, w, h]
|
||||||
|
token_means = ti_heatmaps.mean(dim=[2,3])
|
||||||
|
token_attention_scores = token_means.var(dim=1)
|
||||||
|
|
||||||
|
# Avoid large attention scores in general:
|
||||||
|
reg_loss_0 = 5.0 * torch.stack(att_L2_losses).mean()
|
||||||
|
# Avoid large attention scores for ti tokens, inside the masked region:
|
||||||
|
reg_loss_1 = 1.0 * (torch.relu(ti_heatmaps * ti_masks)**2).mean()
|
||||||
|
# Avoid large attention scores for ti tokens, outside of the masked region:
|
||||||
|
reg_loss_2 = 2.0 * (torch.relu(ti_heatmaps * (1 - ti_masks) + 10)**2).mean()
|
||||||
|
# Make the Ti tokens have similar avg attention scores (equal distribution of concept information over tokens):
|
||||||
|
reg_loss_3 = 1.0 * token_attention_scores.mean()
|
||||||
|
|
||||||
|
if verbose:
|
||||||
|
print(f"reg_loss_0: {reg_loss_0.item():.4f}")
|
||||||
|
print(f"reg_loss_1: {reg_loss_1.item():.4f}")
|
||||||
|
print(f"reg_loss_2: {reg_loss_2.item():.4f}")
|
||||||
|
print(f"reg_loss_3: {reg_loss_3.item():.4f}")
|
||||||
|
|
||||||
|
return (reg_loss_0 + reg_loss_1 + reg_loss_2 + reg_loss_3).to(masks.dtype)
|
||||||
|
|
||||||
|
|
||||||
def compute_snr(noise_scheduler, timesteps):
|
def compute_snr(noise_scheduler, timesteps):
|
||||||
"""
|
"""
|
||||||
@@ -105,14 +179,7 @@ class ConditioningRegularizer:
|
|||||||
def __init__(self, config, embedding_handler):
|
def __init__(self, config, embedding_handler):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.embedding_handler = embedding_handler
|
self.embedding_handler = embedding_handler
|
||||||
self.target_norms = {
|
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
|
||||||
"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.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
|
||||||
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
|
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
|
||||||
|
|
||||||
@@ -122,15 +189,11 @@ class ConditioningRegularizer:
|
|||||||
if tokenizer is None:
|
if tokenizer is None:
|
||||||
idx += 1
|
idx += 1
|
||||||
continue
|
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)
|
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
|
||||||
idx += 1
|
idx += 1
|
||||||
|
|
||||||
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.003, pipe=None):
|
def apply_regularization(self, loss, losses, prompt_embeds_norms, prompt_embeds, std_loss_w = 0.01, pipe=None):
|
||||||
noise_sigma = 0.0
|
noise_sigma = 0.0
|
||||||
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
|
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
|
||||||
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
|
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
|
||||||
|
|||||||
+32
-54
@@ -4,50 +4,30 @@ import subprocess
|
|||||||
import torch
|
import torch
|
||||||
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
|
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||||
|
|
||||||
############################################################################################################
|
def load_models(pretrained_model, device, weight_dtype = torch.float16):
|
||||||
|
|
||||||
SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
|
|
||||||
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
|
|
||||||
|
|
||||||
#SDXL_MODEL_CACHE = "./models/Juggernaut-X-RunDiffusion-NSFW.safetensors"
|
|
||||||
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-X-v10/resolve/main/Juggernaut-X-RunDiffusion-NSFW.safetensors"
|
|
||||||
|
|
||||||
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
|
|
||||||
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
|
|
||||||
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:
|
# check if the model is already downloaded:
|
||||||
if not os.path.exists(pretrained_model['path']):
|
if not os.path.exists(pretrained_model['path']):
|
||||||
download_weights(pretrained_model['url'], pretrained_model['path'])
|
download_weights(pretrained_model['url'], pretrained_model['path'])
|
||||||
|
|
||||||
print(f"Loading model weights from {pretrained_model['path']} with dtype: {weight_dtype}...")
|
tokenizer_two, text_encoder_two = None, None
|
||||||
|
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
|
||||||
|
|
||||||
if pretrained_model['version'] == "sd15":
|
try:
|
||||||
pipe = StableDiffusionPipeline.from_single_file(
|
print("Loading as SDXL model...")
|
||||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
|
||||||
else:
|
|
||||||
pipe = StableDiffusionXLPipeline.from_single_file(
|
pipe = StableDiffusionXLPipeline.from_single_file(
|
||||||
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||||
|
sd_model_version = "sdxl"
|
||||||
|
tokenizer_two = pipe.tokenizer_2
|
||||||
|
text_encoder_two = pipe.text_encoder_2
|
||||||
|
text_encoder_two.requires_grad_(False)
|
||||||
|
text_encoder_two.to(device, dtype=weight_dtype)
|
||||||
|
except:
|
||||||
|
print("Loading as SD15 model...")
|
||||||
|
pipe = StableDiffusionPipeline.from_single_file(
|
||||||
|
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
|
||||||
|
sd_model_version = "sd15"
|
||||||
|
|
||||||
|
print(f"Loaded {sd_model_version} model!")
|
||||||
pipe = pipe.to(device, dtype=weight_dtype)
|
pipe = pipe.to(device, dtype=weight_dtype)
|
||||||
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
|
noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config)
|
||||||
|
|
||||||
@@ -57,24 +37,11 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
|
|||||||
text_encoder_one = pipe.text_encoder
|
text_encoder_one = pipe.text_encoder
|
||||||
|
|
||||||
vae.requires_grad_(False)
|
vae.requires_grad_(False)
|
||||||
if keep_vae_float32:
|
vae.to(device, dtype=weight_dtype)
|
||||||
vae.to(device, dtype=torch.float32)
|
|
||||||
else:
|
|
||||||
vae.to(device, dtype=weight_dtype)
|
|
||||||
if weight_dtype != torch.float32:
|
|
||||||
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but might not be for training..")
|
|
||||||
|
|
||||||
unet.to(device, dtype=weight_dtype)
|
unet.to(device, dtype=weight_dtype)
|
||||||
text_encoder_one.requires_grad_(False)
|
text_encoder_one.requires_grad_(False)
|
||||||
text_encoder_one.to(device, dtype=weight_dtype)
|
text_encoder_one.to(device, dtype=weight_dtype)
|
||||||
|
|
||||||
tokenizer_two = text_encoder_two = None
|
|
||||||
if pretrained_model['version'] == "sdxl":
|
|
||||||
tokenizer_two = pipe.tokenizer_2
|
|
||||||
text_encoder_two = pipe.text_encoder_2
|
|
||||||
text_encoder_two.requires_grad_(False)
|
|
||||||
text_encoder_two.to(device, dtype=weight_dtype)
|
|
||||||
|
|
||||||
return (
|
return (
|
||||||
pipe,
|
pipe,
|
||||||
tokenizer_one,
|
tokenizer_one,
|
||||||
@@ -84,7 +51,7 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
|
|||||||
text_encoder_two,
|
text_encoder_two,
|
||||||
vae,
|
vae,
|
||||||
unet,
|
unet,
|
||||||
)
|
), sd_model_version
|
||||||
|
|
||||||
def download_weights(url, dest):
|
def download_weights(url, dest):
|
||||||
start = time.time()
|
start = time.time()
|
||||||
@@ -108,16 +75,27 @@ def download_weights(url, dest):
|
|||||||
print(f"Downloading {url} took {time.time() - start} seconds")
|
print(f"Downloading {url} took {time.time() - start} seconds")
|
||||||
|
|
||||||
|
|
||||||
def print_trainable_parameters(model, model_name = ''):
|
def print_trainable_parameters(model, model_name=''):
|
||||||
trainable_params = 0
|
trainable_params = 0
|
||||||
all_param = 0
|
all_param = 0
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
all_param += param.numel()
|
all_param += param.numel()
|
||||||
if param.requires_grad and "token_embedding" not in name:
|
if param.requires_grad and "token_embedding" not in name:
|
||||||
trainable_params += param.numel()
|
trainable_params += param.numel()
|
||||||
|
|
||||||
|
def format_param_count(count):
|
||||||
|
if count < 1000:
|
||||||
|
return f"{count}"
|
||||||
|
elif count < 1_000_000:
|
||||||
|
return f"{count/1000:.1f}K"
|
||||||
|
else:
|
||||||
|
return f"{count/1_000_000:.1f}M"
|
||||||
|
|
||||||
line_delimiter = "#" * 80
|
line_delimiter = "#" * 80
|
||||||
print(line_delimiter)
|
print(line_delimiter)
|
||||||
print(
|
print(
|
||||||
f"Trainable {model_name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
|
f"Trainable {model_name} params: {format_param_count(trainable_params)} "
|
||||||
|
f"|| All params: {format_param_count(all_param)} "
|
||||||
|
f"|| trainable = {100 * trainable_params / all_param:.2f}%"
|
||||||
)
|
)
|
||||||
print(line_delimiter)
|
print(line_delimiter)
|
||||||
+42
-8
@@ -3,11 +3,6 @@ import torch
|
|||||||
import prodigyopt
|
import prodigyopt
|
||||||
from typing import Iterable
|
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(
|
def get_unet_optimizer(
|
||||||
prodigy_d_coef: float,
|
prodigy_d_coef: float,
|
||||||
prodigy_growth_factor: float,
|
prodigy_growth_factor: float,
|
||||||
@@ -18,9 +13,12 @@ def get_unet_optimizer(
|
|||||||
):
|
):
|
||||||
## unet_trainable_params can be unet.parameters() or a list of lora params
|
## unet_trainable_params can be unet.parameters() or a list of lora params
|
||||||
|
|
||||||
|
# These learning rates will get overwritten in main.py:
|
||||||
if optimizer_name == "adamw":
|
if optimizer_name == "adamw":
|
||||||
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
|
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
|
||||||
|
elif optimizer_name == "AdamW8bit":
|
||||||
|
import bitsandbytes as bnb
|
||||||
|
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
|
||||||
elif optimizer_name == "prodigy":
|
elif optimizer_name == "prodigy":
|
||||||
# Note: the specific settings of Prodigy seem to matter A LOT
|
# Note: the specific settings of Prodigy seem to matter A LOT
|
||||||
optimizer_unet = prodigyopt.Prodigy(
|
optimizer_unet = prodigyopt.Prodigy(
|
||||||
@@ -40,6 +38,39 @@ def get_unet_optimizer(
|
|||||||
print(f"Created {optimizer_name} optimizer for unet!")
|
print(f"Created {optimizer_name} optimizer for unet!")
|
||||||
return optimizer_unet
|
return optimizer_unet
|
||||||
|
|
||||||
|
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
|
||||||
|
def is_belong_to_blocks(key, blocks):
|
||||||
|
try:
|
||||||
|
for g in blocks:
|
||||||
|
if g in key:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
except Exception as e:
|
||||||
|
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
|
||||||
|
|
||||||
|
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
|
||||||
|
if use_blora:
|
||||||
|
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
|
||||||
|
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
|
||||||
|
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
|
||||||
|
try:
|
||||||
|
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
|
||||||
|
|
||||||
|
attns = [
|
||||||
|
attn_processor_name.rsplit(".", 1)[0]
|
||||||
|
for attn_processor_name, _ in unet.attn_processors.items()
|
||||||
|
if is_belong_to_blocks(attn_processor_name, blocks)
|
||||||
|
]
|
||||||
|
|
||||||
|
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
|
||||||
|
return target_modules
|
||||||
|
except Exception as e:
|
||||||
|
raise type(e)(
|
||||||
|
f"failed to get_target_modules, due to: {e}. "
|
||||||
|
f"Please check the modules specified in --lora_unet_blocks are correct"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_unet_lora_parameters(
|
def get_unet_lora_parameters(
|
||||||
lora_rank,
|
lora_rank,
|
||||||
lora_alpha_multiplier: float,
|
lora_alpha_multiplier: float,
|
||||||
@@ -48,12 +79,15 @@ def get_unet_lora_parameters(
|
|||||||
unet,
|
unet,
|
||||||
pipe,
|
pipe,
|
||||||
):
|
):
|
||||||
|
|
||||||
|
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
|
||||||
|
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
|
||||||
|
|
||||||
unet_lora_config = LoraConfig(
|
unet_lora_config = LoraConfig(
|
||||||
r=lora_rank,
|
r=lora_rank,
|
||||||
lora_alpha=lora_rank * lora_alpha_multiplier,
|
lora_alpha=lora_rank * lora_alpha_multiplier,
|
||||||
init_lora_weights="gaussian",
|
init_lora_weights="gaussian",
|
||||||
target_modules=["to_k", "to_q", "to_v", "to_out.0", "conv2"],
|
target_modules=target_modules,
|
||||||
#target_modules=["conv1", "conv2", "norm1", "norm2", "proj_in"], # TODO grid-search params for sd15
|
|
||||||
use_dora=use_dora,
|
use_dora=use_dora,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+137
-102
@@ -1,7 +1,3 @@
|
|||||||
# Have SwinIR upsample
|
|
||||||
# Have BLIP auto caption
|
|
||||||
# Have CLIPSeg auto mask concept
|
|
||||||
|
|
||||||
import gc
|
import gc
|
||||||
import fnmatch
|
import fnmatch
|
||||||
import mimetypes
|
import mimetypes
|
||||||
@@ -25,6 +21,7 @@ import numpy as np
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from transformers import (
|
from transformers import (
|
||||||
BlipForConditionalGeneration,
|
BlipForConditionalGeneration,
|
||||||
Blip2ForConditionalGeneration,
|
Blip2ForConditionalGeneration,
|
||||||
@@ -38,13 +35,13 @@ from transformers import (
|
|||||||
|
|
||||||
from trainer.utils.io import download_and_prep_training_data
|
from trainer.utils.io import download_and_prep_training_data
|
||||||
from trainer.utils.utils import fix_prompt
|
from trainer.utils.utils import fix_prompt
|
||||||
|
from trainer.config import model_paths
|
||||||
|
|
||||||
import re
|
import re
|
||||||
import openai
|
import openai
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
load_dotenv()
|
load_dotenv()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
|
||||||
client = OpenAI(api_key=OPENAI_API_KEY)
|
client = OpenAI(api_key=OPENAI_API_KEY)
|
||||||
@@ -54,11 +51,9 @@ except:
|
|||||||
client = None
|
client = None
|
||||||
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
|
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
|
||||||
|
|
||||||
MODEL_PATH = "./cache"
|
|
||||||
|
|
||||||
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...)
|
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...)
|
||||||
MIN_GPT_PROMPTS = 3
|
MIN_GPT_PROMPTS = 3
|
||||||
MAX_GPT_PROMPTS = 50
|
MAX_GPT_PROMPTS = 80
|
||||||
|
|
||||||
def _find_files(pattern, dir="."):
|
def _find_files(pattern, dir="."):
|
||||||
"""Return list of files matching pattern in a given directory, in absolute format.
|
"""Return list of files matching pattern in a given directory, in absolute format.
|
||||||
@@ -139,7 +134,7 @@ def swin_ir_sr(
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
model = Swin2SRForImageSuperResolution.from_pretrained(
|
model = Swin2SRForImageSuperResolution.from_pretrained(
|
||||||
model_id, cache_dir=MODEL_PATH
|
model_id, cache_dir = model_paths.get_path("SR")
|
||||||
).to(device)
|
).to(device)
|
||||||
processor = Swin2SRImageProcessor()
|
processor = Swin2SRImageProcessor()
|
||||||
|
|
||||||
@@ -193,9 +188,9 @@ def clipseg_mask_generator(
|
|||||||
|
|
||||||
model = None
|
model = None
|
||||||
if any(target_prompts):
|
if any(target_prompts):
|
||||||
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
|
||||||
model = CLIPSegForImageSegmentation.from_pretrained(
|
model = CLIPSegForImageSegmentation.from_pretrained(
|
||||||
model_id, cache_dir=MODEL_PATH
|
model_id, cache_dir = model_paths.get_path("CLIP")
|
||||||
).to(device)
|
).to(device)
|
||||||
|
|
||||||
masks = []
|
masks = []
|
||||||
@@ -236,7 +231,6 @@ def clipseg_mask_generator(
|
|||||||
|
|
||||||
return masks
|
return masks
|
||||||
|
|
||||||
|
|
||||||
import textwrap
|
import textwrap
|
||||||
def cleanup_prompts_with_chatgpt(
|
def cleanup_prompts_with_chatgpt(
|
||||||
prompts,
|
prompts,
|
||||||
@@ -337,12 +331,12 @@ def extract_gpt_concept_description(gpt_completion, concept_mode):
|
|||||||
return concept_name
|
return concept_name
|
||||||
|
|
||||||
|
|
||||||
def post_process_captions(captions, text, concept_mode, job_seed):
|
def post_process_captions(captions, text, concept_mode, job_seed, skip_gpt_cleanup=False):
|
||||||
text = text.strip()
|
text = text.strip()
|
||||||
gpt_cleanup_worked = False
|
gpt_cleanup_worked = False
|
||||||
gpt_concept_description = None
|
gpt_concept_description = None
|
||||||
|
|
||||||
if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client:
|
if (len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client) and not skip_gpt_cleanup:
|
||||||
retry_count = 0
|
retry_count = 0
|
||||||
while retry_count < 5:
|
while retry_count < 5:
|
||||||
try:
|
try:
|
||||||
@@ -408,14 +402,14 @@ def blip_caption_dataset(
|
|||||||
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|
||||||
if "blip2" in model_id:
|
if "blip2" in model_id:
|
||||||
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
|
||||||
model = Blip2ForConditionalGeneration.from_pretrained(
|
model = Blip2ForConditionalGeneration.from_pretrained(
|
||||||
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
|
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
|
||||||
).to(device)
|
).to(device)
|
||||||
else:
|
else:
|
||||||
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
|
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
|
||||||
model = BlipForConditionalGeneration.from_pretrained(
|
model = BlipForConditionalGeneration.from_pretrained(
|
||||||
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
|
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
|
||||||
).to(device)
|
).to(device)
|
||||||
|
|
||||||
for i, image in enumerate(tqdm(images)):
|
for i, image in enumerate(tqdm(images)):
|
||||||
@@ -424,6 +418,7 @@ def blip_caption_dataset(
|
|||||||
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65)
|
out = model.generate(**inputs, max_length=100, do_sample=True, top_k=40, temperature=0.65)
|
||||||
captions[i] = processor.decode(out[0], skip_special_tokens=True)
|
captions[i] = processor.decode(out[0], skip_special_tokens=True)
|
||||||
|
|
||||||
|
model.to("cpu")
|
||||||
del model
|
del model
|
||||||
gc.collect()
|
gc.collect()
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
@@ -445,51 +440,6 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
|
|||||||
os.remove(output_path)
|
os.remove(output_path)
|
||||||
return base64_image
|
return base64_image
|
||||||
|
|
||||||
def gpt4_v_get_description(config, images):
|
|
||||||
if config.concept_mode == "object":
|
|
||||||
description = "object"
|
|
||||||
prompt = "Give a concise visual descriptioni of the object/figure/thing that all the grid-images have in common with at most 10 words. Dont start with statements like 'The image features...', just describe what you see."
|
|
||||||
|
|
||||||
elif config.concept_mode == "face":
|
|
||||||
description = "face"
|
|
||||||
prompt = "All the grid images depict a single person. Visually describe this person with at most 10 words. Dont start with statements like 'The image features...', just describe what you see. (eg an asian woman with long black hair)"
|
|
||||||
|
|
||||||
elif config.concept_mode == "style":
|
|
||||||
description = ""
|
|
||||||
prompt = "All these images share a common aesthetic style. Describe this style with at most 7 words. Dont start with statements like 'The image features...', just describe what you see. (eg impressionism collage surrealism)"
|
|
||||||
|
|
||||||
if not OPENAI_API_KEY:
|
|
||||||
print(f"Skipping GPT-4 Vision description because OPENAI_API_KEY is not set.")
|
|
||||||
return description
|
|
||||||
|
|
||||||
|
|
||||||
headers = {
|
|
||||||
"Content-Type": "application/json",
|
|
||||||
"Authorization": f"Bearer {OPENAI_API_KEY}"
|
|
||||||
}
|
|
||||||
|
|
||||||
# TODO sample a grid img:
|
|
||||||
# .... TODO
|
|
||||||
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
|
|
||||||
|
|
||||||
payload = {
|
|
||||||
"model": "gpt-4-turbo",
|
|
||||||
"messages": [
|
|
||||||
{
|
|
||||||
"role": "user",
|
|
||||||
"content": [
|
|
||||||
{"type": "text", "text": prompt},
|
|
||||||
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{base64_image}", "detail": "high"}}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"max_tokens": 60
|
|
||||||
}
|
|
||||||
|
|
||||||
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
|
|
||||||
answer = response.json()["choices"][0]["message"]["content"]
|
|
||||||
return captions
|
|
||||||
|
|
||||||
def gpt4_v_caption_dataset(
|
def gpt4_v_caption_dataset(
|
||||||
images, captions,
|
images, captions,
|
||||||
batch_size=4,
|
batch_size=4,
|
||||||
@@ -500,6 +450,7 @@ def gpt4_v_caption_dataset(
|
|||||||
return captions
|
return captions
|
||||||
|
|
||||||
prompt = "Concisely describe this image without assumptions with at most 20 words. Dont start with statements like 'The image features...', just describe what you see."
|
prompt = "Concisely describe this image without assumptions with at most 20 words. Dont start with statements like 'The image features...', just describe what you see."
|
||||||
|
#prompt = "Concisely describe the main subject(s) in the image without assumptions with at most 20 words. Ignore the background and only describe the people / animals / objects in the foreground. Dont start with statements like 'The image features...', just describe what you see."
|
||||||
|
|
||||||
headers = {
|
headers = {
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
@@ -510,7 +461,7 @@ def gpt4_v_caption_dataset(
|
|||||||
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
|
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
|
||||||
|
|
||||||
payload = {
|
payload = {
|
||||||
"model": "gpt-4-turbo",
|
"model": "gpt-4o",
|
||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
@@ -548,17 +499,84 @@ def gpt4_v_caption_dataset(
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
device = "cuda:0" if torch.cuda.is_available() else "cpu"
|
||||||
|
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
|
||||||
|
@torch.no_grad()
|
||||||
|
def florence_caption_dataset(images, captions):
|
||||||
|
|
||||||
|
#workaround for unnecessary flash_attn requirement
|
||||||
|
from unittest.mock import patch
|
||||||
|
from transformers.dynamic_module_utils import get_imports
|
||||||
|
from transformers import AutoProcessor, AutoModelForCausalLM
|
||||||
|
|
||||||
|
def fixed_get_imports(filename: str | os.PathLike) -> list[str]:
|
||||||
|
if not str(filename).endswith("modeling_florence2.py"):
|
||||||
|
return get_imports(filename)
|
||||||
|
imports = get_imports(filename)
|
||||||
|
try:
|
||||||
|
imports.remove("flash_attn")
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
return imports
|
||||||
|
|
||||||
|
with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement
|
||||||
|
model = AutoModelForCausalLM.from_pretrained("microsoft/Florence-2-large", attn_implementation="sdpa", device_map=device, torch_dtype=torch_dtype,trust_remote_code=True)
|
||||||
|
|
||||||
|
processor = AutoProcessor.from_pretrained("microsoft/Florence-2-large", trust_remote_code=True, cache_dir = model_paths.get_path("FLORENCE"))
|
||||||
|
|
||||||
|
for i, image in enumerate(tqdm(images)):
|
||||||
|
if captions[i] is None:
|
||||||
|
#prompt = random.choice(["<CAPTION>", "<DETAILED_CAPTION>"])
|
||||||
|
prompt = "<CAPTION>"
|
||||||
|
prompt = "<DETAILED_CAPTION>"
|
||||||
|
prompt = "<MORE_DETAILED_CAPTION>"
|
||||||
|
|
||||||
|
inputs = processor(text=prompt, images=image, return_tensors="pt").to(device, torch_dtype)
|
||||||
|
generated_ids = model.generate(
|
||||||
|
input_ids=inputs["input_ids"],
|
||||||
|
pixel_values=inputs["pixel_values"],
|
||||||
|
max_new_tokens=1024,
|
||||||
|
num_beams=random.choice([2,3,4])
|
||||||
|
)
|
||||||
|
|
||||||
|
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
|
||||||
|
parsed_answer = processor.post_process_generation(generated_text, task=prompt, image_size=(image.width, image.height))
|
||||||
|
caption = parsed_answer[prompt]
|
||||||
|
captions[i] = caption.replace("The image shows a ", "A ")
|
||||||
|
|
||||||
|
model.to('cpu')
|
||||||
|
del model
|
||||||
|
del processor
|
||||||
|
gc.collect()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
return captions
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def caption_dataset(
|
def caption_dataset(
|
||||||
images: List[Image.Image],
|
images: List[Image.Image],
|
||||||
captions: List[str],
|
captions: List[str],
|
||||||
caption_model: Literal[str] = "blip"
|
caption_model: Literal["blip", "gpt4-v", "florence"] = "blip"
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
|
|
||||||
|
# if all captions are already generated, we dont need to do anything:
|
||||||
|
if all(captions):
|
||||||
|
print(f"All captions loaded from disk, skipping captioning...")
|
||||||
|
return captions
|
||||||
|
|
||||||
if "blip" in caption_model:
|
if "blip" in caption_model:
|
||||||
captions = blip_caption_dataset(images, captions)
|
captions = blip_caption_dataset(images, captions)
|
||||||
elif "gpt4-v" in caption_model:
|
elif "gpt4-v" in caption_model:
|
||||||
captions = gpt4_v_caption_dataset(images, captions)
|
captions = gpt4_v_caption_dataset(images, captions)
|
||||||
|
elif "florence" in caption_model:
|
||||||
|
captions = florence_caption_dataset(images, captions)
|
||||||
|
else:
|
||||||
|
print("WARNING: not using any captions!")
|
||||||
|
captions = [""] * len(images)
|
||||||
|
|
||||||
|
gc.collect()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
return captions
|
return captions
|
||||||
|
|
||||||
@@ -629,20 +647,44 @@ def random_crop(image, scale=(0.85, 0.95)):
|
|||||||
|
|
||||||
return image.crop((left, top, left + new_width, top + new_height))
|
return image.crop((left, top, left + new_width, top + new_height))
|
||||||
|
|
||||||
def gaussian_blur(image):
|
def gaussian_blur(image, radius = 1.0):
|
||||||
return image.filter(ImageFilter.GaussianBlur(radius=1))
|
return image.filter(ImageFilter.GaussianBlur(radius=radius))
|
||||||
|
|
||||||
def augment_image(image):
|
def augment_image(image):
|
||||||
image = hue_augmentation(image)
|
image = hue_augmentation(image)
|
||||||
image = color_jitter(image)
|
image = color_jitter(image)
|
||||||
image = random_crop(image)
|
image = random_crop(image)
|
||||||
if random.random() < 0.5:
|
if random.random() < 0.5:
|
||||||
image = gaussian_blur(image)
|
image = gaussian_blur(image, radius = random.uniform(0.0, 1.0))
|
||||||
return image
|
return image
|
||||||
|
|
||||||
def round_to_nearest_multiple(x, multiple):
|
def round_to_nearest_multiple(x, multiple):
|
||||||
return int(float(multiple) * round(float(x) / float(multiple)))
|
return int(float(multiple) * round(float(x) / float(multiple)))
|
||||||
|
|
||||||
|
'''
|
||||||
|
For Stable Diffusion 1.5, outputs are optimised around 512x512 pixels. Many common fine-tuned versions of SD1.5 are optimised around 768x768. The best resolutions for common aspect ratios are typically:
|
||||||
|
1:1 (square): 512x512, 768x768
|
||||||
|
3:2 (landscape): 768x512
|
||||||
|
2:3 (portrait): 512x768
|
||||||
|
4:3 (landscape): 768x576
|
||||||
|
3:4 (portrait): 576x768
|
||||||
|
16:9 (widescreen): 912x512
|
||||||
|
9:16 (tall): 512x912
|
||||||
|
|
||||||
|
For SDXL, outputs are optimised around 1024x1024 pixels. The best resolutions for common aspect ratios are typically:
|
||||||
|
stable-diffusion-xl-1024-v0-9 supports generating images at the following dimensions:
|
||||||
|
1024 x 1024
|
||||||
|
1152 x 896
|
||||||
|
896 x 1152
|
||||||
|
1216 x 832
|
||||||
|
832 x 1216
|
||||||
|
1344 x 768
|
||||||
|
768 x 1344
|
||||||
|
1536 x 640
|
||||||
|
640 x 1536
|
||||||
|
|
||||||
|
'''
|
||||||
|
|
||||||
def calculate_new_dimensions(target_size, target_aspect_ratio):
|
def calculate_new_dimensions(target_size, target_aspect_ratio):
|
||||||
"""
|
"""
|
||||||
Calculate the new width and height given a target size and aspect ratio.
|
Calculate the new width and height given a target size and aspect ratio.
|
||||||
@@ -661,8 +703,6 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
|
|||||||
return [new_width, new_height]
|
return [new_width, new_height]
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def load_and_save_masks_and_captions(
|
def load_and_save_masks_and_captions(
|
||||||
config,
|
config,
|
||||||
concept_mode: str,
|
concept_mode: str,
|
||||||
@@ -703,9 +743,10 @@ def load_and_save_masks_and_captions(
|
|||||||
n_length = len(files)
|
n_length = len(files)
|
||||||
files = sorted(files)[:n_length]
|
files = sorted(files)[:n_length]
|
||||||
|
|
||||||
images, captions = [], []
|
images, captions, img_paths = [], [], []
|
||||||
for file in files:
|
for file in files:
|
||||||
images.append(load_image_with_orientation(file))
|
images.append(load_image_with_orientation(file))
|
||||||
|
img_paths.append(file)
|
||||||
caption_file = os.path.splitext(file)[0] + ".txt"
|
caption_file = os.path.splitext(file)[0] + ".txt"
|
||||||
if os.path.exists(caption_file) and use_dataset_captions:
|
if os.path.exists(caption_file) and use_dataset_captions:
|
||||||
with open(caption_file, "r") as f:
|
with open(caption_file, "r") as f:
|
||||||
@@ -746,41 +787,39 @@ def load_and_save_masks_and_captions(
|
|||||||
upscale_margin = 0.75
|
upscale_margin = 0.75
|
||||||
images = swin_ir_sr(images, target_size=(int(config.train_img_size[0]*upscale_margin), int(config.train_img_size[0]*upscale_margin)))
|
images = swin_ir_sr(images, target_size=(int(config.train_img_size[0]*upscale_margin), int(config.train_img_size[0]*upscale_margin)))
|
||||||
|
|
||||||
if add_lr_flips and len(images) < 40:
|
if add_lr_flips and len(images) < MAX_GPT_PROMPTS:
|
||||||
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
|
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
|
||||||
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
|
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
|
||||||
captions = captions + captions
|
captions = captions + captions
|
||||||
|
|
||||||
# It's nice if we can achieve the gpt pass, so pre-augment the images if there's very few:
|
|
||||||
# Ensure we have at least 'augment_imgs_up_to_n' images through augmentation
|
|
||||||
aug_imgs, aug_caps = [],[]
|
|
||||||
# if we still have a very small amount of imgs, do some basic augmentation:
|
|
||||||
while len(images) + len(aug_imgs) < MIN_GPT_PROMPTS:
|
|
||||||
print(f"Adding augmented version of each training img...")
|
|
||||||
aug_imgs.extend([augment_image(image) for image in images])
|
|
||||||
aug_caps.extend(captions)
|
|
||||||
|
|
||||||
images.extend(aug_imgs)
|
print(f"Generating {len(images)} captions using {caption_model} in {concept_mode} mode...")
|
||||||
captions.extend(aug_caps)
|
captions = caption_dataset(images, captions, caption_model = caption_model)
|
||||||
|
|
||||||
|
# Save captions back to disk:
|
||||||
|
for i, img_path in enumerate(img_paths):
|
||||||
|
caption_path = os.path.splitext(img_path)[0] + ".txt"
|
||||||
|
with open(caption_path, "w") as f:
|
||||||
|
f.write(captions[i])
|
||||||
|
|
||||||
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
|
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
|
||||||
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
|
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
|
||||||
images = images[:MAX_GPT_PROMPTS-1]
|
images = images[:MAX_GPT_PROMPTS-1]
|
||||||
captions = captions[:MAX_GPT_PROMPTS-1]
|
captions = captions[:MAX_GPT_PROMPTS-1]
|
||||||
|
|
||||||
if len(images) > 50 and caption_model != "blip":
|
|
||||||
print(f"Captioning a lot of ({len(images)}) images --> falling back to using blip!")
|
|
||||||
caption_model = "blip"
|
|
||||||
|
|
||||||
print(f"Generating {len(images)} captions using mode: {concept_mode}...")
|
|
||||||
captions = caption_dataset(images, captions, caption_model = caption_model)
|
|
||||||
|
|
||||||
# Cleanup prompts using chatgpt:
|
# Cleanup prompts using chatgpt:
|
||||||
captions = [fix_prompt(caption) for caption in captions]
|
captions = [fix_prompt(caption) for caption in captions]
|
||||||
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
|
trigger_text = ""
|
||||||
|
gpt_concept_description = None
|
||||||
|
if not config.disable_ti:
|
||||||
|
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed, skip_gpt_cleanup=config.skip_gpt_cleanup)
|
||||||
|
|
||||||
|
if config.prompt_modifier:
|
||||||
|
print(config.prompt_modifier)
|
||||||
|
captions = [config.prompt_modifier.format(caption) for caption in captions]
|
||||||
|
|
||||||
aug_imgs, aug_caps = [],[]
|
aug_imgs, aug_caps = [],[]
|
||||||
# if we still have a very small amount of imgs, do some basic augmentation:
|
# if we still have a small amount of imgs, do some basic augmentation:
|
||||||
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
|
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
|
||||||
print(f"Adding augmented version of each training img...")
|
print(f"Adding augmented version of each training img...")
|
||||||
aug_imgs.extend([augment_image(image) for image in images])
|
aug_imgs.extend([augment_image(image) for image in images])
|
||||||
@@ -789,19 +828,18 @@ def load_and_save_masks_and_captions(
|
|||||||
images.extend(aug_imgs)
|
images.extend(aug_imgs)
|
||||||
captions.extend(aug_caps)
|
captions.extend(aug_caps)
|
||||||
|
|
||||||
|
|
||||||
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
|
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
|
||||||
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
|
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
|
||||||
mask_target_prompts = gpt_concept_description
|
mask_target_prompts = gpt_concept_description
|
||||||
|
|
||||||
if mask_target_prompts is None or config.concept_mode == "style":
|
if mask_target_prompts is None or config.concept_mode == "style":
|
||||||
print("Disabling CLIP-segmentation")
|
|
||||||
mask_target_prompts = ""
|
mask_target_prompts = ""
|
||||||
temp = 999
|
temp = 999
|
||||||
else:
|
else:
|
||||||
temp = config.clipseg_temperature
|
temp = config.clipseg_temperature
|
||||||
|
|
||||||
print(f"Generating {len(images)} masks...")
|
print(f"Generating {len(images)} masks...")
|
||||||
|
|
||||||
# Make sure we have a bias for the background pixels to never 100% ignore them
|
# Make sure we have a bias for the background pixels to never 100% ignore them
|
||||||
background_bias = 0.05
|
background_bias = 0.05
|
||||||
if not use_face_detection_instead:
|
if not use_face_detection_instead:
|
||||||
@@ -854,16 +892,11 @@ def load_and_save_masks_and_captions(
|
|||||||
os.remove(os.path.join(output_dir, file))
|
os.remove(os.path.join(output_dir, file))
|
||||||
|
|
||||||
os.makedirs(output_dir, exist_ok=True)
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
|
||||||
# Make sure we've correctly inserted the TOK into every caption:
|
if config.disable_ti:
|
||||||
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
|
|
||||||
for caption in captions:
|
|
||||||
print(caption)
|
|
||||||
|
|
||||||
if config.remove_ti_token_from_prompts:
|
|
||||||
print('------------------ WARNING -------------------')
|
print('------------------ WARNING -------------------')
|
||||||
print("Removing 'TOK, ' from captions...")
|
print("Removing 'TOK, ' from captions...")
|
||||||
print("This will completely break textual_inversion!!")
|
print("This will completely disable textual_inversion!!")
|
||||||
print('------------------ WARNING -------------------')
|
print('------------------ WARNING -------------------')
|
||||||
if gpt_concept_description:
|
if gpt_concept_description:
|
||||||
replace_str = gpt_concept_description
|
replace_str = gpt_concept_description
|
||||||
@@ -871,6 +904,8 @@ def load_and_save_masks_and_captions(
|
|||||||
replace_str = ""
|
replace_str = ""
|
||||||
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
|
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
|
||||||
captions = [caption.replace("TOK", replace_str) for caption in captions]
|
captions = [caption.replace("TOK", replace_str) for caption in captions]
|
||||||
|
else:
|
||||||
|
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
|
||||||
|
|
||||||
# iterate through the images, masks, and captions and add a row to the dataframe for each
|
# iterate through the images, masks, and captions and add a row to the dataframe for each
|
||||||
print("Saving final training dataset...")
|
print("Saving final training dataset...")
|
||||||
|
|||||||
@@ -0,0 +1,364 @@
|
|||||||
|
from functools import reduce
|
||||||
|
from diffusers import StableDiffusionXLPipeline
|
||||||
|
from diffusers.models.attention_processor import AttnProcessor2_0, Attention
|
||||||
|
from typing import Optional
|
||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from diffusers.utils.deprecation_utils import deprecate
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import math
|
||||||
|
from einops.layers.torch import Reduce
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from mpl_toolkits.axes_grid1 import make_axes_locatable
|
||||||
|
|
||||||
|
def plot_token_attention_loss(folder, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio):
|
||||||
|
try:
|
||||||
|
batch_index = 0
|
||||||
|
timestep = timesteps[batch_index].item()
|
||||||
|
if timestep > 700:
|
||||||
|
batch_index += 1
|
||||||
|
timestep = timesteps[batch_index].item()
|
||||||
|
|
||||||
|
folder = os.path.join(folder, "attention_heatmaps")
|
||||||
|
os.makedirs(folder, exist_ok=True)
|
||||||
|
|
||||||
|
token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])]
|
||||||
|
plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens
|
||||||
|
|
||||||
|
# Calculate global min and max for consistent colormap
|
||||||
|
all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float()
|
||||||
|
for i in plot_token_indices]
|
||||||
|
vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps])
|
||||||
|
|
||||||
|
# Heatmap plots
|
||||||
|
fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices),
|
||||||
|
figsize=(3 * len(plot_token_indices), 10))
|
||||||
|
title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}"
|
||||||
|
fig.suptitle(title_str, fontsize=16)
|
||||||
|
|
||||||
|
for idx, text_token_index in enumerate(plot_token_indices):
|
||||||
|
heatmap = all_heatmaps[idx]
|
||||||
|
im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax)
|
||||||
|
axes[idx].set_title(f"{token_strings[text_token_index]}")
|
||||||
|
axes[idx].axis("off")
|
||||||
|
|
||||||
|
# Add colorbar
|
||||||
|
divider = make_axes_locatable(axes[idx])
|
||||||
|
cax = divider.append_axes("bottom", size="5%", pad=0.05)
|
||||||
|
plt.colorbar(im, cax=cax, orientation="horizontal")
|
||||||
|
|
||||||
|
plt.tight_layout()
|
||||||
|
fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||||
|
plt.close(fig)
|
||||||
|
|
||||||
|
# Histogram plot
|
||||||
|
fig, ax = plt.subplots(figsize=(12, 6))
|
||||||
|
ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16)
|
||||||
|
ax.set_xlabel("Attention Value", fontsize=12)
|
||||||
|
ax.set_ylabel("Frequency", fontsize=12)
|
||||||
|
|
||||||
|
for idx, text_token_index in enumerate(plot_token_indices):
|
||||||
|
heatmap = all_heatmaps[idx]
|
||||||
|
ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index],
|
||||||
|
alpha=0.5, density=True)
|
||||||
|
|
||||||
|
ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10)
|
||||||
|
ax.grid(alpha=0.3)
|
||||||
|
plt.tight_layout()
|
||||||
|
|
||||||
|
# Add token attention loss as text
|
||||||
|
plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}",
|
||||||
|
transform=ax.transAxes, ha='right', va='top',
|
||||||
|
bbox=dict(facecolor='white', edgecolor='black', alpha=0.8))
|
||||||
|
|
||||||
|
plt.ylim(0, 0.6)
|
||||||
|
fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight')
|
||||||
|
plt.close(fig)
|
||||||
|
except:
|
||||||
|
print("Failed to plot token attention loss")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# Find all instances of AttnProcessor2_0 in the UNet
|
||||||
|
def find_attnprocessor2_0(unet):
|
||||||
|
|
||||||
|
module_names = []
|
||||||
|
"""
|
||||||
|
this function assumes that there are fewer than 50 down blocks, attention modules and transformer blocks
|
||||||
|
if you're not sure, feel free to set it to an arbitrarily large number
|
||||||
|
don't worry, it won't slow anything down.
|
||||||
|
"""
|
||||||
|
|
||||||
|
for block_type in ["down_blocks", "up_blocks"]:
|
||||||
|
for down_block_index in range(50):
|
||||||
|
for attentions_index in range(50):
|
||||||
|
for transformer_blocks_index in range(50):
|
||||||
|
example_module_name = f"{block_type}.{down_block_index}.attentions.{attentions_index}.transformer_blocks.{transformer_blocks_index}.attn2.processor"
|
||||||
|
|
||||||
|
try:
|
||||||
|
module = get_module_by_name(module=unet, name = example_module_name)
|
||||||
|
assert isinstance(module, AttnProcessor2_0), f"Expected module to be an instance of AttnProcessor2_0 but found it to be: {type(module)}"
|
||||||
|
# print(f"Found: {example_module_name}")
|
||||||
|
module_names.append(example_module_name)
|
||||||
|
except AttributeError:
|
||||||
|
# print(f"Ignored name: {example_module_name}\nsince it does not exist")
|
||||||
|
pass
|
||||||
|
print(f"Found: {len(module_names)} modules")
|
||||||
|
return module_names
|
||||||
|
|
||||||
|
class DAAMLossAttnProcessor2_0:
|
||||||
|
r"""
|
||||||
|
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, name: str):
|
||||||
|
if not hasattr(F, "scaled_dot_product_attention"):
|
||||||
|
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||||
|
|
||||||
|
self.name = name
|
||||||
|
self.cross_attention_scores = None
|
||||||
|
self.reduce_op = Reduce(
|
||||||
|
"batch heads img text -> batch img text",
|
||||||
|
reduction="sum"
|
||||||
|
)
|
||||||
|
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
attn: Attention,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
temb: Optional[torch.Tensor] = None,
|
||||||
|
*args,
|
||||||
|
**kwargs,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if len(args) > 0 or kwargs.get("scale", None) is not None:
|
||||||
|
deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`."
|
||||||
|
deprecate("scale", "1.0.0", deprecation_message)
|
||||||
|
|
||||||
|
residual = hidden_states
|
||||||
|
if attn.spatial_norm is not None:
|
||||||
|
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||||
|
|
||||||
|
input_ndim = hidden_states.ndim
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
batch_size, channel, height, width = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||||
|
|
||||||
|
batch_size, sequence_length, _ = (
|
||||||
|
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||||
|
)
|
||||||
|
|
||||||
|
if attention_mask is not None:
|
||||||
|
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||||
|
# scaled_dot_product_attention expects attention_mask shape to be
|
||||||
|
# (batch, heads, source_length, target_length)
|
||||||
|
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||||
|
|
||||||
|
if attn.group_norm is not None:
|
||||||
|
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||||
|
|
||||||
|
query = attn.to_q(hidden_states)
|
||||||
|
|
||||||
|
"""
|
||||||
|
Mayukh's experiment
|
||||||
|
"""
|
||||||
|
mayukh_experiment = False
|
||||||
|
if encoder_hidden_states is not None:
|
||||||
|
"""
|
||||||
|
this triggers cross attn
|
||||||
|
"""
|
||||||
|
mayukh_experiment = True
|
||||||
|
|
||||||
|
if encoder_hidden_states is None:
|
||||||
|
encoder_hidden_states = hidden_states
|
||||||
|
elif attn.norm_cross:
|
||||||
|
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||||
|
|
||||||
|
key = attn.to_k(encoder_hidden_states)
|
||||||
|
value = attn.to_v(encoder_hidden_states)
|
||||||
|
|
||||||
|
inner_dim = key.shape[-1]
|
||||||
|
head_dim = inner_dim // attn.heads
|
||||||
|
|
||||||
|
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||||
|
|
||||||
|
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||||
|
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||||
|
|
||||||
|
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||||
|
# TODO: add support for attn.scale when we move to Torch 2.1
|
||||||
|
hidden_states = F.scaled_dot_product_attention(
|
||||||
|
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||||
|
)
|
||||||
|
|
||||||
|
if mayukh_experiment:
|
||||||
|
# Calculate QK^T
|
||||||
|
qk_t = torch.matmul(query, key.transpose(-2, -1))
|
||||||
|
|
||||||
|
# Calculate attention scores (scaled QK^T)
|
||||||
|
d_k = query.size(-1) # Assuming the last dimension is the embedding dimension
|
||||||
|
attention_scores = qk_t / math.sqrt(d_k)
|
||||||
|
|
||||||
|
attention_scores = self.reduce_op(
|
||||||
|
attention_scores,
|
||||||
|
)
|
||||||
|
self.cross_attention_scores = attention_scores
|
||||||
|
|
||||||
|
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||||
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
|
||||||
|
# linear proj
|
||||||
|
hidden_states = attn.to_out[0](hidden_states)
|
||||||
|
# dropout
|
||||||
|
hidden_states = attn.to_out[1](hidden_states)
|
||||||
|
|
||||||
|
if input_ndim == 4:
|
||||||
|
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||||
|
|
||||||
|
if attn.residual_connection:
|
||||||
|
hidden_states = hidden_states + residual
|
||||||
|
|
||||||
|
hidden_states = hidden_states / attn.rescale_output_factor
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
class DAAMLoss:
|
||||||
|
def __init__(self, attention_processors: list[DAAMLossAttnProcessor2_0]):
|
||||||
|
self.attention_processors = attention_processors
|
||||||
|
self.layer_names = [
|
||||||
|
x.name for x in attention_processors
|
||||||
|
]
|
||||||
|
|
||||||
|
def process_and_stack_attention_scores(self, img_ratio: float):
|
||||||
|
reshaped_tensors = []
|
||||||
|
min_heatmap_pixels = np.inf
|
||||||
|
|
||||||
|
# Process each attention score
|
||||||
|
for processor in self.attention_processors:
|
||||||
|
score = processor.cross_attention_scores
|
||||||
|
bs, seq_len, channels = score.shape
|
||||||
|
|
||||||
|
# Calculate width and height based on img_ratio
|
||||||
|
width = round(math.sqrt(seq_len * img_ratio))
|
||||||
|
height = round(width / img_ratio)
|
||||||
|
reshaped_score = rearrange(score, 'b (h w) c -> b h w c', h=height, w=width)
|
||||||
|
reshaped_tensors.append(reshaped_score)
|
||||||
|
|
||||||
|
if height*width < min_heatmap_pixels:
|
||||||
|
min_heatmap_pixels = height*width
|
||||||
|
min_heatmap_shape = height, width
|
||||||
|
|
||||||
|
# Interpolate and standardize all tensors to the same size
|
||||||
|
for i, heatmap in enumerate(reshaped_tensors):
|
||||||
|
if heatmap.shape[1] * heatmap.shape[2] != min_heatmap_pixels:
|
||||||
|
# Interpolating to match the smallest tensor size uniformly
|
||||||
|
heatmap = F.interpolate(heatmap.permute(0, 3, 1, 2), size=(min_heatmap_shape[0], min_heatmap_shape[1]), mode='bicubic').permute(0, 2, 3, 1)
|
||||||
|
reshaped_tensors[i] = heatmap
|
||||||
|
|
||||||
|
|
||||||
|
# Stack all tensors along the first dimension
|
||||||
|
stacked_tensor = torch.stack(reshaped_tensors, dim=0)
|
||||||
|
return stacked_tensor
|
||||||
|
|
||||||
|
def get_all_cross_attention_scores(self):
|
||||||
|
cross_attention_scores = {}
|
||||||
|
|
||||||
|
for p in self.attention_processors:
|
||||||
|
cross_attention_scores[
|
||||||
|
p.name
|
||||||
|
] = p.cross_attention_scores
|
||||||
|
|
||||||
|
return cross_attention_scores
|
||||||
|
|
||||||
|
def get_image_heatmap(self, text_token_index: int, layer_name: str, img_ratio: float):
|
||||||
|
cross_attention_scores = self.get_all_cross_attention_scores()
|
||||||
|
assert layer_name in list(cross_attention_scores.keys())
|
||||||
|
|
||||||
|
cross_attention_scores_single_token = cross_attention_scores[layer_name][:,:,text_token_index]
|
||||||
|
|
||||||
|
width = round(math.sqrt(cross_attention_scores_single_token.shape[-1] * img_ratio))
|
||||||
|
height = round(width / img_ratio)
|
||||||
|
|
||||||
|
heatmap = rearrange(
|
||||||
|
cross_attention_scores_single_token,
|
||||||
|
"batch (height width) -> batch height width",
|
||||||
|
height = height,
|
||||||
|
width = width
|
||||||
|
)
|
||||||
|
|
||||||
|
return heatmap
|
||||||
|
|
||||||
|
def get_the_daam_heatmap(self, text_token_index: int, img_ratio: float, resize = 'min'):
|
||||||
|
all_heatmaps = []
|
||||||
|
for layer_name in self.layer_names:
|
||||||
|
heatmap = self.get_image_heatmap(
|
||||||
|
text_token_index=text_token_index,
|
||||||
|
layer_name=layer_name,
|
||||||
|
img_ratio = img_ratio
|
||||||
|
)
|
||||||
|
all_heatmaps.append(heatmap)
|
||||||
|
|
||||||
|
if resize == 'max':
|
||||||
|
heatmap_height = max(heatmap.shape[1] for heatmap in all_heatmaps)
|
||||||
|
heatmap_width = max(heatmap.shape[2] for heatmap in all_heatmaps)
|
||||||
|
elif resize == 'min':
|
||||||
|
heatmap_height = min(heatmap.shape[1] for heatmap in all_heatmaps)
|
||||||
|
heatmap_width = min(heatmap.shape[2] for heatmap in all_heatmaps)
|
||||||
|
|
||||||
|
## now resize all_heatmaps to (batch, heatmap_height, heatmap_width) using F.interpolate
|
||||||
|
resized_heatmaps = [
|
||||||
|
F.interpolate(input = x.unsqueeze(1), size = (heatmap_height, heatmap_width)).squeeze(1)
|
||||||
|
for x in all_heatmaps
|
||||||
|
]
|
||||||
|
|
||||||
|
resized_heatmaps = torch.stack(resized_heatmaps)
|
||||||
|
|
||||||
|
## Average the heatmaps:
|
||||||
|
avg_heatmap = resized_heatmaps.mean(dim=0)
|
||||||
|
|
||||||
|
return avg_heatmap
|
||||||
|
|
||||||
|
def get_module_by_name(module: nn.Module, name: str):
|
||||||
|
"""Retrieve a module nested in another by its access string."""
|
||||||
|
if name == "":
|
||||||
|
return module
|
||||||
|
names = name.split(sep=".")
|
||||||
|
return reduce(getattr, names, module)
|
||||||
|
|
||||||
|
|
||||||
|
def init_daam_loss(pipeline):
|
||||||
|
## find out where the attention processor thingies are
|
||||||
|
module_names = find_attnprocessor2_0(
|
||||||
|
unet = pipeline.unet
|
||||||
|
)
|
||||||
|
|
||||||
|
all_daam_attention_processors = []
|
||||||
|
# override the attention processor thingies
|
||||||
|
for name in module_names:
|
||||||
|
# Get parent module and attribute name
|
||||||
|
parent_name = ".".join(name.split(".")[:-1])
|
||||||
|
attr_name = name.split(".")[-1]
|
||||||
|
|
||||||
|
# Get the parent module
|
||||||
|
parent_module = get_module_by_name(module=pipeline.unet, name=parent_name)
|
||||||
|
|
||||||
|
daam_attention_processor = DAAMLossAttnProcessor2_0(name=name)
|
||||||
|
all_daam_attention_processors.append(daam_attention_processor)
|
||||||
|
# Set the attribute
|
||||||
|
setattr(parent_module, attr_name, daam_attention_processor)
|
||||||
|
|
||||||
|
# Verify the replacement
|
||||||
|
current_module = get_module_by_name(module=pipeline.unet, name=name)
|
||||||
|
assert isinstance(current_module, DAAMLossAttnProcessor2_0)
|
||||||
|
|
||||||
|
daam_loss = DAAMLoss(
|
||||||
|
attention_processors=all_daam_attention_processors
|
||||||
|
)
|
||||||
|
return pipeline, daam_loss
|
||||||
@@ -290,6 +290,18 @@ def load_image_with_orientation(path, mode="RGB"):
|
|||||||
elif orientation == 8:
|
elif orientation == 8:
|
||||||
image = image.rotate(90, expand=True)
|
image = image.rotate(90, expand=True)
|
||||||
|
|
||||||
|
if image.mode == 'P':
|
||||||
|
image = image.convert('RGBA')
|
||||||
|
if image.mode == 'CMYK':
|
||||||
|
image = image.convert('RGB')
|
||||||
|
|
||||||
|
# Remove alpha channel if present
|
||||||
|
if image.mode in ('RGBA', 'LA'):
|
||||||
|
background = Image.new('RGB', image.size, (255, 255, 255))
|
||||||
|
background.paste(image, mask=image.split()[3]) # 3 is the alpha channel
|
||||||
|
image = background
|
||||||
|
|
||||||
|
# Convert to the desired mode
|
||||||
return image.convert(mode)
|
return image.convert(mode)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+13
-4
@@ -100,15 +100,17 @@ def print_system_info():
|
|||||||
|
|
||||||
# Print disk space information
|
# Print disk space information
|
||||||
disk_usage = psutil.disk_usage('/')
|
disk_usage = psutil.disk_usage('/')
|
||||||
free_disk = disk_usage.free // (1024 * 1024)
|
total_disk = disk_usage.total // (1024 * 1024)
|
||||||
|
used_disk = disk_usage.used // (1024 * 1024)
|
||||||
percent_disk_used = disk_usage.percent
|
percent_disk_used = disk_usage.percent
|
||||||
print(f"Free disk space: {free_disk} MB with {percent_disk_used}% used")
|
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
|
||||||
|
|
||||||
# Print RAM information
|
# Print RAM information
|
||||||
virtual_mem = psutil.virtual_memory()
|
virtual_mem = psutil.virtual_memory()
|
||||||
|
total_ram = virtual_mem.total // (1024 * 1024)
|
||||||
current_ram = virtual_mem.used // (1024 * 1024)
|
current_ram = virtual_mem.used // (1024 * 1024)
|
||||||
percent_ram_used = virtual_mem.percent
|
percent_ram_used = virtual_mem.percent
|
||||||
print(f"Current used RAM: {current_ram} MB with {percent_ram_used}% used")
|
print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f'Error in gathering system info: {str(e)}')
|
print(f'Error in gathering system info: {str(e)}')
|
||||||
@@ -122,6 +124,13 @@ def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1
|
|||||||
|
|
||||||
# Flatten and concatenate all parameters into a single tensor
|
# Flatten and concatenate all parameters into a single tensor
|
||||||
all_params = torch.cat([p.data.view(-1) for p in parameters])
|
all_params = torch.cat([p.data.view(-1) for p in parameters])
|
||||||
|
|
||||||
|
# count number of parameters:
|
||||||
|
n_params = len(all_params)
|
||||||
|
|
||||||
|
if n_params == 0 or n_params > 1e9:
|
||||||
|
return
|
||||||
|
|
||||||
norm = torch.norm(all_params)
|
norm = torch.norm(all_params)
|
||||||
|
|
||||||
# Convert to CPU for plotting
|
# Convert to CPU for plotting
|
||||||
@@ -228,7 +237,7 @@ def plot_token_stds(token_std_dict, save_path='token_stds.png', target_value_dic
|
|||||||
|
|
||||||
from scipy.signal import savgol_filter
|
from scipy.signal import savgol_filter
|
||||||
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'):
|
def plot_loss(loss_dict, save_path='losses.png', window_length=31, polyorder=3, default_color='gray'):
|
||||||
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red'}
|
colormap = {'img_loss': 'blue', 'tot_loss': 'green', 'covariance_tok_reg_loss': 'orange', 'concept_description_loss': 'red', 'token_attention_loss': 'purple'}
|
||||||
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
|
values_to_add_to_title = ['concept_description_loss', 'covariance_tok_reg_loss']
|
||||||
plot_smoothed = ['img_loss']
|
plot_smoothed = ['img_loss']
|
||||||
|
|
||||||
|
|||||||
@@ -2,9 +2,9 @@
|
|||||||
val_prompts = {}
|
val_prompts = {}
|
||||||
val_prompts['style'] = [
|
val_prompts['style'] = [
|
||||||
'a beautiful mountainous landscape, boulders, fresh water stream, setting sun',
|
'a beautiful mountainous landscape, boulders, fresh water stream, setting sun',
|
||||||
'the stunning skyline of New York City',
|
'the stunning skyline of New York City, setting sun, skyscrapers, wallpaper',
|
||||||
'fruit hanging from a tree, highly detailed texture, soil, rain, drops, photo realistic, surrealism, highly detailed, 8k macrophotography',
|
'fruit hanging from a tree, highly detailed texture, soil, rain, drops, photo realistic, surrealism, highly detailed, 8k macrophotography',
|
||||||
'the Taj Mahal, stunning wallpaper',
|
'the Taj Mahal, stunning wallpaper, architecture, ancient, marble, white, intricate, detailed',
|
||||||
'A majestic tree rooted in circuits, leaves shimmering with data streams, stands as a beacon where the digital dawn caresses the fog-laden, binary soil—a symphony of pixels and chlorophyll.',
|
'A majestic tree rooted in circuits, leaves shimmering with data streams, stands as a beacon where the digital dawn caresses the fog-laden, binary soil—a symphony of pixels and chlorophyll.',
|
||||||
'A beautiful octopus, with swirling tendrils and a pulsating heart of fiery opal hues, hovers ethereally against a starry void, sculpted through a meticulous flame-working technique.',
|
'A beautiful octopus, with swirling tendrils and a pulsating heart of fiery opal hues, hovers ethereally against a starry void, sculpted through a meticulous flame-working technique.',
|
||||||
'a stunning image of an aston martin sportscar',
|
'a stunning image of an aston martin sportscar',
|
||||||
|
|||||||
@@ -1,31 +0,0 @@
|
|||||||
{
|
|
||||||
"output_dir": "lora_models/xander_sd15_final",
|
|
||||||
"sd_model_version": "sd15",
|
|
||||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
|
|
||||||
"concept_mode": "face",
|
|
||||||
"seed": 0,
|
|
||||||
"resolution": 512,
|
|
||||||
"train_batch_size": 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
|
|
||||||
}
|
|
||||||
@@ -1,25 +0,0 @@
|
|||||||
{
|
|
||||||
"output_dir": "lora_models/object",
|
|
||||||
"sd_model_version": "sdxl",
|
|
||||||
"lora_training_urls": "/home/rednax/Documents/datasets/DOV/lizzo/full body",
|
|
||||||
"concept_mode": "object",
|
|
||||||
"seed": 1,
|
|
||||||
"resolution": 640,
|
|
||||||
"train_batch_size": 4,
|
|
||||||
"n_sample_imgs": 4,
|
|
||||||
"max_train_steps": 420,
|
|
||||||
"token_warmup_steps": 0,
|
|
||||||
"checkpointing_steps": 60,
|
|
||||||
"gradient_accumulation_steps": 1,
|
|
||||||
"n_tokens": 2,
|
|
||||||
"ti_lr": 0.001,
|
|
||||||
"ti_weight_decay": 0.0005,
|
|
||||||
"text_encoder_lora_optimizer": null,
|
|
||||||
"text_encoder_lora_lr": 1.0e-4,
|
|
||||||
"text_encoder_lora_weight_decay": 1e-5,
|
|
||||||
"text_encoder_lora_rank": 12,
|
|
||||||
"lora_rank": 12,
|
|
||||||
"use_dora": false,
|
|
||||||
"caption_model": "gpt4-v",
|
|
||||||
"debug": true
|
|
||||||
}
|
|
||||||
@@ -1,29 +0,0 @@
|
|||||||
{
|
|
||||||
"output_dir": "lora_models/does_best",
|
|
||||||
"sd_model_version": "sd15",
|
|
||||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip",
|
|
||||||
"concept_mode": "style",
|
|
||||||
"seed": 1,
|
|
||||||
"resolution": 640,
|
|
||||||
"train_batch_size": 4,
|
|
||||||
"n_sample_imgs": 6,
|
|
||||||
"max_train_steps": 600,
|
|
||||||
"token_warmup_steps": 0,
|
|
||||||
"checkpointing_steps": 100,
|
|
||||||
"gradient_accumulation_steps": 1,
|
|
||||||
"n_tokens": 2,
|
|
||||||
"ti_lr": 0.001,
|
|
||||||
"ti_weight_decay": 0.0005,
|
|
||||||
|
|
||||||
"remove_ti_token_from_prompts": false,
|
|
||||||
"text_encoder_lora_optimizer": null,
|
|
||||||
"text_encoder_lora_lr": 1.0e-4,
|
|
||||||
"text_encoder_lora_weight_decay": 1e-5,
|
|
||||||
"text_encoder_lora_rank": 12,
|
|
||||||
|
|
||||||
"unet_lr": 0.001,
|
|
||||||
"lora_rank": 16,
|
|
||||||
"use_dora": false,
|
|
||||||
"caption_model": "blip",
|
|
||||||
"debug": true
|
|
||||||
}
|
|
||||||
@@ -1,29 +0,0 @@
|
|||||||
{
|
|
||||||
"output_dir": "lora_models/Journey",
|
|
||||||
"sd_model_version": "sdxl",
|
|
||||||
"lora_training_urls": "/home/rednax/Documents/datasets/journey",
|
|
||||||
"concept_mode": "style",
|
|
||||||
"seed": 0,
|
|
||||||
"resolution": 512,
|
|
||||||
"train_batch_size": 4,
|
|
||||||
"n_sample_imgs": 6,
|
|
||||||
"max_train_steps": 1000,
|
|
||||||
"token_warmup_steps": 0,
|
|
||||||
"checkpointing_steps": 100,
|
|
||||||
"gradient_accumulation_steps": 1,
|
|
||||||
"n_tokens": 2,
|
|
||||||
"ti_lr": 0.001,
|
|
||||||
"ti_weight_decay": 0.0005,
|
|
||||||
"text_encoder_lora_optimizer": null,
|
|
||||||
"text_encoder_lora_lr": 1.0e-4,
|
|
||||||
"text_encoder_lora_weight_decay": 1e-5,
|
|
||||||
"text_encoder_lora_rank": 12,
|
|
||||||
|
|
||||||
"unet_lr": 0.001,
|
|
||||||
"prodigy_d_coef": 1.0,
|
|
||||||
"unet_prodigy_growth_factor": 1.05,
|
|
||||||
"lora_rank": 16,
|
|
||||||
"use_dora": true,
|
|
||||||
"caption_model": "gpt4-v",
|
|
||||||
"debug": true
|
|
||||||
}
|
|
||||||
Reference in New Issue
Block a user