37 Commits
Author SHA1 Message Date
aiXander 8b340a6edc unpushed changes 2024-03-29 10:38:52 -07:00
aiXander 322b3da61e merge 2024-03-15 09:12:01 -07:00
aiXander 7da3e64b1e updates 2024-03-15 09:10:58 -07:00
mayukhdeb 5088ad3754 train_pti: switch to 500 steps + fix precision arg 2024-03-15 02:04:30 -07:00
mayukhdeb d5e56dc6e1 use keyword args for sanity 2024-03-15 02:04:02 -07:00
mayukhdeb 4563279eca config: forbid extras 2024-03-15 02:03:08 -07:00
mayukhdeb bf892826f7 oopsie 2024-03-15 00:48:47 -07:00
mayukhdeb 661f41c7c6 hardcode l1_penalty to be 0.05 if concept_mode == "style" 2024-03-15 00:48:14 -07:00
mayukhdeb fad22f61cc trainer_pti: keep all args in preprocess and init 2024-03-15 00:41:28 -07:00
mayukhdeb bbdd8d4762 trainer: fix prodigy_d_coef 2024-03-15 00:20:11 -07:00
mayukhdeb 42adc67c3a fix prodigy_d_coef default value 2024-03-15 00:19:43 -07:00
mayukhdeb f020bcf218 trainer: fix undefined variable 2024-03-14 22:04:34 -07:00
aiXander 838856fcef Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-03-14 21:32:25 -07:00
mayukhdeb 87fc89141b switch to fp32 2024-03-14 21:22:04 -07:00
aiXander 12bd221dbe trying to fix trainer 2024-03-14 21:05:46 -07:00
aiXander 8d2517741c more updates 2024-03-14 20:39:19 -07:00
aiXander 580d7867a8 make training work again 2024-03-14 20:12:30 -07:00
aiXander b35f4529dd fix more bugs 2024-03-14 19:45:35 -07:00
aiXander 28a9772cf6 large amounts of bugfixes 2024-03-14 19:24:24 -07:00
aiXander 0722ab6ce7 sync training args, add some bugfixes 2024-03-14 18:18:52 -07:00
aiXander 416a4622c2 add download weights 2024-03-14 17:45:56 -07:00
Xander Steenbrugge 40edbd0e85 add special_params.json and training_args.json 2024-03-14 17:05:26 -07:00
mayukhdeb 2cd456020b add concept_mode arg 2024-03-14 12:51:06 -07:00
mayukhdeb d4b550792d render_images regardless of debug or not 2024-03-14 12:50:42 -07:00
mayukhdeb 01d112f140 cleanup 2024-03-14 07:04:51 -07:00
mayukhdeb f87ab80048 save train args even for intermediate checkpoints 2024-03-14 06:59:29 -07:00
mayukhdeb 5ebc20ace0 return only output save dir 2024-03-14 06:57:36 -07:00
mayukhdeb 6b2f19ceb9 dont store validation prompts 2024-03-14 06:50:44 -07:00
mayukhdeb 53bda05d3d dont store validation_promts in train config 2024-03-14 06:50:16 -07:00
mayukhdeb b4a76f12ba stop ignoring trainer folder + remove train.py 2024-03-14 06:42:58 -07:00
mayukhdeb 9b2bc5da0f add trainer class 2024-03-14 06:42:58 -07:00
mayukhdeb 3387701f9a add config 2024-03-14 06:42:58 -07:00
mayukhdeb c9c9c89442 add utils 2024-03-14 06:42:58 -07:00
mayukhdeb 9de4a47c31 add init files 2024-03-14 06:42:58 -07:00
mayukhdeb c8bf5ce3b5 add prepare_prompt_for_lora 2024-03-14 06:42:58 -07:00
mayukhdeb a4ba97aebc move stuff into trainer module 2024-03-14 06:42:58 -07:00
mayukhdeb a03e1f6cb0 move files 2024-03-14 06:42:58 -07:00
64 changed files with 2833 additions and 8161 deletions
-25
View File
@@ -1,25 +0,0 @@
# The .dockerignore file excludes files from the container build process.
# https://docs.docker.com/engine/reference/builder/#dockerignore-file
# Exclude Git files
.git
.github
.gitignore
# Exclude Python cache files
__pycache__
.mypy_cache
.pytest_cache
.ruff_cache
# exclude trained model rars:
*.rar
# Dev folders:
debug
xander
datasets
rendered_images
# trained models:
lora_models/*
-26
View File
@@ -1,26 +0,0 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'edenartlab' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+6 -18
View File
@@ -1,25 +1,13 @@
models
lora_models
remove
cache
__pycache__
.ipynb_checkpoints/
models
lora_models*
eden_lora_training_runs/
datasets
*.tar
.env
.cog
xander*.sh
.huggingface
rendered_images*
gridsearch*
aesthetic_score_best_model.pth
# experiment folders:
scripts/plots
conditioning_spaces/
training_args_x_*.json
xander_configs/
tests/
debug/*
!debug/*.py
File diff suppressed because it is too large Load Diff
@@ -1,782 +0,0 @@
{
"last_node_id": 26,
"last_link_id": 49,
"nodes": [
{
"id": 15,
"type": "Note Plus (mtb)",
"pos": {
"0": 301,
"1": -120,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
519.1570279742359,
181.43570138536967
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"## How to trigger the LoRa:\nEden's trainer trains a token by default (if not disabled) which triggers the concept.\n\nFor \"object\" and \"face\" mode you just refer to your concept directly with the token, eg: \"a photo of embedding:MY\\_NAME\\_embedding\"\n\nFor \"style\" mode, you just prepend \"in the style of embedding:MY\\_NAME\\_embedding\" in the beginning of your prompt!",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 13,
"type": "Note Plus (mtb)",
"pos": {
"0": -506,
"1": -142,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
684.8613808966753,
209.8075855491436
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"## How to use:\n\n1. Find the training folder in ComfyUI/outputs\n2. Go into the checkpoints folder\n3. Pick the best checkpoint by looking at the validation_grid\n4. Copy the ***_embeddings.safetensors file to ComfyUI/models/embeddings\n5. Copy the ***_LoRa.safetensors file to ComfyUI/models/loras\n6. Hit Refresh in your ComfyUI\n7. Adjust this workflow to load both of those!\n8. Tweak the lora strength and embedding token strength to get the best results!\n\n\nHave fun! :)",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1162,
188
],
"size": {
"0": 140,
"1": 46
},
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 8
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
9
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-451,
162
],
"size": [
429.578383119695,
98
],
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
10
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
12
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
8
],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"zavychromaxl_v90.safetensors"
]
},
{
"id": 9,
"type": "SaveImage",
"pos": [
1327,
188
],
"size": {
"0": 453.85968017578125,
"1": 407.6841125488281
},
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 9
}
],
"properties": {
"Node name for S&R": "SaveImage"
},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
859,
26
],
"size": [
271.419023112996,
106
],
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
1024,
1024,
1
]
},
{
"id": 10,
"type": "LoraLoader",
"pos": [
31,
163
],
"size": {
"0": 337.44464111328125,
"1": 126
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 10
},
{
"name": "clip",
"type": "CLIP",
"link": 12
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
24
],
"shape": 3,
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
25,
26
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LoraLoader"
},
"widgets_values": [
"Eden_Token_LoRa_sdxl_LoRa.safetensors",
0.7000000000000001,
0.7000000000000001
]
},
{
"id": 3,
"type": "KSampler",
"pos": [
863,
186
],
"size": [
267.5352926614355,
262
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 24
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 40
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 44
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
7
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
1093427037772868,
"randomize",
30,
8,
"euler_ancestral",
"normal",
1
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
402,
124
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 25
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
43
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"(in the style of embedding:Eden_Token_LoRa_sdxl_embeddings:1.0), \n\nwhat do i say to make me exist, oriental mythical beasts, in the golden danish age, in the history of television in the style of light violet and light red, serge najjar, playful and whimsical, associated press photo, afrofuturism-inspired, alasdair mclellan, electronic media"
]
},
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 26
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
45
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
""
]
},
{
"id": 22,
"type": "ControlNetApplyAdvanced",
"pos": [
672,
741
],
"size": {
"0": 315,
"1": 166
},
"flags": {},
"order": 11,
"mode": 4,
"inputs": [
{
"name": "positive",
"type": "CONDITIONING",
"link": 43
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 45
},
{
"name": "control_net",
"type": "CONTROL_NET",
"link": 38
},
{
"name": "image",
"type": "IMAGE",
"link": 47
}
],
"outputs": [
{
"name": "positive",
"type": "CONDITIONING",
"links": [
40
],
"shape": 3,
"slot_index": 0
},
{
"name": "negative",
"type": "CONDITIONING",
"links": [
44
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "ControlNetApplyAdvanced"
},
"widgets_values": [
1,
0,
1
]
},
{
"id": 23,
"type": "ControlNetLoader",
"pos": [
255,
733
],
"size": {
"0": 315,
"1": 58
},
"flags": {},
"order": 4,
"mode": 4,
"outputs": [
{
"name": "CONTROL_NET",
"type": "CONTROL_NET",
"links": [
38
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ControlNetLoader"
},
"widgets_values": [
"SDXL/controlnet-canny-sdxl-1.0/diffusion_pytorch_model_V2.safetensors"
]
},
{
"id": 25,
"type": "AIO_Preprocessor",
"pos": [
258,
857
],
"size": {
"0": 315,
"1": 82
},
"flags": {},
"order": 7,
"mode": 4,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 48
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
47,
49
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "AIO_Preprocessor"
},
"widgets_values": [
"CannyEdgePreprocessor",
512
]
},
{
"id": 26,
"type": "PreviewImage",
"pos": [
679,
958
],
"size": [
307.39388136076934,
34.317686638573264
],
"flags": {},
"order": 10,
"mode": 4,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 49
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 24,
"type": "LoadImage",
"pos": [
-112,
857
],
"size": [
315,
314
],
"flags": {},
"order": 5,
"mode": 4,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
48
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"000050.jpg",
"image"
]
}
],
"links": [
[
2,
5,
0,
3,
3,
"LATENT"
],
[
7,
3,
0,
8,
0,
"LATENT"
],
[
8,
4,
2,
8,
1,
"VAE"
],
[
9,
8,
0,
9,
0,
"IMAGE"
],
[
10,
4,
0,
10,
0,
"MODEL"
],
[
12,
4,
1,
10,
1,
"CLIP"
],
[
24,
10,
0,
3,
0,
"MODEL"
],
[
25,
10,
1,
6,
0,
"CLIP"
],
[
26,
10,
1,
7,
0,
"CLIP"
],
[
38,
23,
0,
22,
2,
"CONTROL_NET"
],
[
40,
22,
0,
3,
1,
"CONDITIONING"
],
[
43,
6,
0,
22,
0,
"CONDITIONING"
],
[
44,
22,
1,
3,
2,
"CONDITIONING"
],
[
45,
7,
0,
22,
1,
"CONDITIONING"
],
[
47,
25,
0,
22,
3,
"IMAGE"
],
[
48,
24,
0,
25,
0,
"IMAGE"
],
[
49,
25,
0,
26,
0,
"IMAGE"
]
],
"groups": [
{
"title": "Optional Controlnet",
"bounding": [
-180,
645,
1205,
536
],
"color": "#3f789e",
"font_size": 24,
"locked": false
}
],
"config": {},
"extra": {
"ds": {
"scale": 0.683013455365071,
"offset": {
"0": 493.3539751601979,
"1": 181.8365096487169
}
}
},
"version": 0.4
}
-312
View File
@@ -1,312 +0,0 @@
{
"last_node_id": 21,
"last_link_id": 35,
"nodes": [
{
"id": 4,
"type": "Display Any (rgthree)",
"pos": [
899,
263
],
"size": {
"0": 349.7635803222656,
"1": 88.69296264648438
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 34,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 3,
"type": "Display Any (rgthree)",
"pos": [
898,
150
],
"size": {
"0": 347.36749267578125,
"1": 87.71485137939453
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 33,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 5,
"type": "Display Any (rgthree)",
"pos": [
897,
373
],
"size": {
"0": 349.3501892089844,
"1": 98.81207275390625
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "source",
"type": "*",
"link": 35,
"dir": 3
}
],
"properties": {
"Node name for S&R": "Display Any (rgthree)"
},
"widgets_values": [
""
]
},
{
"id": 2,
"type": "PreviewImage",
"pos": [
1307,
121
],
"size": {
"0": 557.6970825195312,
"1": 461.9648742675781
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 32
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 18,
"type": "Note Plus (mtb)",
"pos": {
"0": 411,
"1": -278,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": {
"0": 1447.654541015625,
"1": 329.1551513671875
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"\n## [Eden](https://www.eden.art/)\nThis LoRa trainer was made by the team behind https://www.eden.art/, led by https://x.com/xsteenbrugge\n\nIf you make awesome stuff w this trainer, give us a shout at:\nhttps://x.com/eden_art_ or\nhttps://www.instagram.com/eden.art____/\n\n---\n\n---\n\n## SD15 + SDXL:\nThis trainer works for both SD15 and SDXL models, but the default settings are primarily tuned for SDXL models. By selecting a specific ckpt_name, the trainer will automatically know if its an SDXL or SD15 model!\n\n---\n\n### A note on Embeddings:\nThis trainer optionally trains a textual inversion token into the LoRa, this is highly recommended when using SDXL models, but also means you have to load that token embedding when doing inference! See the example workflows in the repo.\n\n\nThe default settings work great for SDXL, SD15 usually need more training steps (eg 800) and sometimes benefits from disabling ti_training.\n\n---\n\n### A note on captioning:\nImages will get automatically captioned. It is recommended to put a .env file in the root of this custom node repo with your OpenAI API key, if found that will trigger a prompt_cleanup function that significantly improves results!\n",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 19,
"type": "Note Plus (mtb)",
"pos": {
"0": 48,
"1": 106,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": [
325.27622360123536,
609
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [],
"title": "Unnamed",
"properties": {},
"widgets_values": [
"\n## A note on settings:\n\nIf you disbale\\_ti (not recommended) you'll get a normal LoRa that does not use a token embedding, in that case you typically need to train for a bit longer and increase the unet_lr.\n\n---\n\n\"training_images\" can both be a path to a local folder or a url to a public .zip file of imgs which will get downloaded.\n\n---\n\nYou can provide custom captions by placing a filename.txt file for each filename.jpg in the training_images folder\n\n---\n\nI highly recommend to keep the training resolution at either 512 or 768.\n\n---\n\nn_tokens = 1 is currently broken, need to fix that.\n\n---\n\nAn embedding + LoRa checkpoint will get saved every **save\\_checkpoint\\_every\\_n\\_steps**. Based on the sample image grid you can then pick the best checkpoint to use in your workflows!\n\n---\n\nSetting debug=True will save a bunch of additional graphs and visualizations to track whats happening during training for advanced users.",
"markdown",
"",
"one_dark"
],
"color": "#432",
"bgcolor": "#653",
"shape": 1
},
{
"id": 21,
"type": "Eden_LoRa_trainer",
"pos": [
413,
124
],
"size": [
412.5926450093166,
591.4499899627035
],
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "sample_images",
"type": "IMAGE",
"links": [
32
],
"shape": 3,
"slot_index": 0
},
{
"name": "lora_path",
"type": "STRING",
"links": [
33
],
"shape": 3,
"slot_index": 1
},
{
"name": "embedding_path",
"type": "STRING",
"links": [
34
],
"shape": 3,
"slot_index": 2
},
{
"name": "final_msg",
"type": "STRING",
"links": [
35
],
"shape": 3,
"slot_index": 3
}
],
"properties": {
"Node name for S&R": "Eden_LoRa_trainer"
},
"widgets_values": [
"https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip",
"style",
"Eden_Token_LoRa",
"zavychromaxl_v90.safetensors",
512,
4,
300,
0.001,
0.0005,
16,
false,
3,
200,
6,
0.7,
false,
40092,
"randomize"
]
}
],
"links": [
[
32,
21,
0,
2,
0,
"IMAGE"
],
[
33,
21,
1,
3,
0,
"*"
],
[
34,
21,
2,
4,
0,
"*"
],
[
35,
21,
3,
5,
0,
"*"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.6830134553650705,
"offset": [
106.80403405562511,
337.0121827057539
]
}
},
"version": 0.4
}
-43
View File
@@ -1,43 +0,0 @@
Open Source Native License (OSNL)
Version 0.1 - March 1, 2024
Preamble
The Open Source Native License (OSNL) is designed to ensure that software remains free and open, fostering innovation and knowledge sharing within the community. It grants individuals, researchers, and commercial entities who open source their primary business assets, the freedom to use the software in any manner they choose.
This distinctive approach aims to balance the benefits of open-source development with the realities of commercial enterprise. It ensures that software remains a shared, community-driven resource while enabling businesses to thrive in an open-source ecosystem. Additional licenses are available for non-open source commercial entities.
1. Definitions
- "This License" refers to Version 1.0 of the Open Source Native License.
- "The Program" refers to the software distributed under this License.
- "You" refers to the individual or entity utilizing or contributing to the Program.
- "Primary Business Assets" are the core resources, capabilities, and technology that constitute the main value proposition and operational basis of your business.
2. Grant of License
Subject to the terms and conditions of this License, you are hereby granted a free, perpetual, worldwide, non-exclusive, no-charge, royalty-free, irrevocable license to use, reproduce, modify, distribute, and sublicense the Program, provided you comply with the following condition:
- Individual or researcher: You are granted the rights to use, modify, distribute, and contribute to the Program for any purpose, including educational, research, and personal projects, without the necessity to make your personal projects open source, provided these activities do not constitute a commercial enterprise. For any use that transitions to commercial purposes, the conditions applicable to commercial entities as outlined in this License will then apply.
- Commercial Entity who meets open source condition: Your primary business assets, including all core technologies, software, and platforms, must be available under an OSI-approved open source license or OSNL. This condition does not apply to ancillary or peripheral services not constituting primary business assets.
2.1 Commercial Use by Non-Open Source Businesses
Non-open source businesses that wish to utilize the Program or its derivatives as a component of their products or services are required to obtain an additional license. These entities must proactively contact Banodoco to request such a license. Banodoco reserves the right, at its own discretion, to grant or deny this additional license. Until an additional license is granted by Banodoco, non-open source businesses are not authorized to exercise any rights provided under this License regarding the use of the Program or its derivatives.
3. Redistribution
You may reproduce and distribute copies of the Program or derivative works thereof in any medium, with or without modifications, provided that you meet the following conditions:
- You must give any recipients of the Program a copy of this License.
- You must ensure that any modified files carry prominent notices stating that you changed the files.
- You must disclose the source of the Program, and if you distribute any portion of it in a compiled or object code form, you must also provide the full source code under this License.
- Any distribution of the Program or derivative works must comply with the Primary Business Open Source Condition.
4. Disclaimer of Warranty
THE PROGRAM IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES, OR OTHER LIABILITY ARISING FROM THE USE OF THE PROGRAM.
5. General
This License does not grant permission to use the trade names, trademarks, service marks, or product names of the Licensor, except as required for reasonable and customary use in describing the origin of the Program.
+42 -64
View File
@@ -1,58 +1,9 @@
# Trainer
This trainer was developed by the [**Eden** team](https://eden.art/), you can try our hosted version of the trainer in [**our app**](https://app.eden.art/).
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
The outputs of this trainer are fully compatible with ComfyUI and AUTO111, see documentation [here](https://docs.eden.art/docs/guides/concepts/#exporting-loras-for-use-in-other-tools).
A full guide on training can be found in [**our docs**](https://docs.eden.art/docs/guides/concepts/#training).
<p align="center">
<strong>Training images:</strong><br>
<img src="assets/xander_training_images.jpg" alt="Image 1" style="width:80%;"/>
</p>
<p align="center">
<strong>Generated imgs with trained LoRa:</strong><br>
<img src="assets/xander_generated_images.jpg" alt="Image 2" style="width:80%;"/>
</p>
### The trainer can be run in 4 different ways:
- [**as a hosted service on our website**](https://app.eden.art/)
- [**as a hosted service through replicate**](https://replicate.com/edenartlab/sdxl-lora-trainer)
- **as a ComfyUI node**
- **as a standalone python script**
### Using in ComfyUI:
- Example workflows for how to run the trainer and do inference with it can be found in `/ComfyUI_workflows`
- Importantly this trainer uses a chatgpt call to cleanup the auto-generated prompts and inject the trainable token, this will only work if you have a .env file containing your OPENAI key in the root of the repo dir that contains a single line: `OPENAI_API_KEY=your_key_string` Everything will work without this, but results will be better if you set this up, especially for 'face' and 'object' modes.
### The trainer supports 3 default modes:
- **style**: used for learning the aesthetic style of a collection of images.
- **face**: used for learning a specific face (can be human, character, ...).
- **object**: will learn a specific object or thing featured in the training images.
<p align="center">
<strong>Style training example:</strong><br>
<img src="assets/style_training_example.jpg" alt="Image 1" style="width:80%;"/>
</p>
Code for finetuning and training LoRa modules on top of Stable Diffusion.
## Setup
Install all dependencies using
`pip install -r requirements.txt`
then you can simply run:
`python main.py train_configs/training_args.json`
to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job.
---
You can also run this through Replicate using cog (~docker image):
1. Install Replicate 'cog':
```
@@ -60,25 +11,52 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
sudo chmod +x /usr/local/bin/cog
```
2. Build the image with `cog build`
3. Run a training run with `sh cog_test_train.sh`
4. You can also go into the container with `cog run /bin/bash`
2. Build the image with `sudo cog build`
3. Run a training run with `sudo sh test_train.sh`
## Full unet finetuning
When running this trainer in native python, you can also perform full unet finetuning using something like (adjust to your needs)
`python main.py train_configs/full_finetuning_example.json`
## TODO's
Bugs:
- pure textual inversion for SD15 does not seem to work well... (but it works amazingly well for SDXL...) ---> if anyone can figure this one out I'd be forever grateful!
- figure out why training is 3x slower through comfyui node versus just running main.py as a python job..?
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
Code / Cleanup:
- turn all/most of the args of the main() function in trainer_pti.py and the preprocess() function into a clean args_dict that makes it easy to add and distribute new parameters over the code and save these args to a .json file at the end.
- Modularize the logic in train.py as much as possible, trying to minimize dev work that needs to happen when SD3 drops
- make a clean train.py entrypoint that can be run as a normal python command (instead of having to use cog)
- make it so the textual_inversion optimizer only optimizes the actual trained token embeddings instead of all of them + resetting later
- test if the trained concepts with peft are compatible with ComfyUI / AUTO1111
Algo:
- Add aspect_ratio bucketing into the dataloader so we can train on non-square images (take this from https://github.com/kohya-ss/sd-scripts)
- test if textual inversion training can also happen with prodigy_optimizer
- the random initialization of the token embeddings has a relatively large impact on the final outcome, there are prob ways to reduce
this random variance, eg CLIP_similarity pretraining.
- Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM
- it looks the like gpu-utilization is only like 65-70% during training: whats the bottleneck? Can we speed this up?
Bugfixing:
see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037
- check how the pipe() objects work under the hood in HF diffusers library, is there a difference w how the unet is called in the training loop / which args it gets?
- Try to find out why the diffusers training script works for sd15 and ours doesnt:
See here: https://huggingface.co/blog/sdxl_lora_advanced_script
and here: https://github.com/huggingface/diffusers/tree/main/examples/advanced_diffusion_training
- figure out how to adaptively set lora_scale at inference time using peft + diffusers? (https://github.com/huggingface/peft/blob/main/src/peft/tuners/lora/layer.py#L240)
Bigger improvements:
- integrate Flux / SD3
- Add multi-concept training (multiple things represented by multiple tokens, trained into a single LoRa)
- add stronger token regularization (eg CelebBasis spanning basis)
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- Add multi-token training
- implement perfusion: https://research.nvidia.com/labs/par/Perfusion/
- implement prompt-aligned: https://prompt-aligned.github.io/
- make compatible with ziplora: https://ziplora.github.io/
Tuning Experiments once code is fully ready:
- test if VAE weight_type actually matters for training
- try-out conditioning noise injection during training to increase robustness
- re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam)
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
but it then recovers. Can we avoid this collapse? Is the learning rate too high?
- gradient_accumulation
- offset noise
- AB test Dora vs Lora
- sweep n_trainable_tokens to inject
-10
View File
@@ -1,10 +0,0 @@
import os
import sys
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from node import Eden_LoRa_trainer
NODE_CLASS_MAPPINGS = {
"Eden_LoRa_trainer": Eden_LoRa_trainer,
}
Binary file not shown.

Before

Width:  |  Height:  |  Size: 430 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 915 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 342 KiB

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