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
54 changed files with 2774 additions and 7183 deletions
-32
View File
@@ -1,32 +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/*
# Ignore the entire models folder by default:
models/*
### Include pipeline models: ###
!models/juggernaut_reborn.safetensors
!models/juggernaut_v6.safetensors
+6 -17
View File
@@ -1,24 +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:
conditioning_spaces/
training_args_x_*.json
xander_configs/
tests/
debug/*
!debug/*.py
+32 -75
View File
@@ -1,42 +1,9 @@
# Trainer
This trainer was developed by the [**Eden** team](https://eden.art/)
It's a highly optimized trainer that can be used for both full finetuning and training LoRa modules on top of Stable Diffusion.
It uses a single training script and loss module that works for both **SDv15** and **SDXL**!
The outputs of this trainer are fully compatible with ComfyUI and AUTO111.
<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 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.
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':
```
@@ -44,62 +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`
## Automatic Checkpoint Evaluation
This script uses CLIP img/txt similarity scores to evaluate how good the LoRa is vs how overfit.
Download the aesthetic predictor model checkpoint first from google drive. This should give you a file named: `aesthetic_score_best_model.pth` (99.2 MB)
```bash
gdown 1thEIlXVc8lkULVUBY9Ab45tsOERxkjxns
```
Once the model is downloaded, you can run the eval script with the following CLI args:
- `output_folder`: this is where the outputs of the model get saved as jpeg files
- `lora_path`: path to your LoRA checkpoint (make sure you edit `path_to_your_model_checkpoints` to point to the correct folder. It generally ends with something like `checkpoint-600` where `600` was the training step)
- `output_json`: save all scores in this json file
- `config_filename`: config file used for training
```bash
python3 evaluate.py \
--output_folder eval_images \
--lora_path path_to_your_model_checkpoint \
--output_json eval_results.json \
--config_filename training_args.json
```
2. Build the image with `sudo cog build`
3. Run a training run with `sudo sh test_train.sh`
## 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!
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:
- Improve some of the chatgpt functionality:
- separate the "gpt_description" / "gpt_segmentation" prompt calls and make them run on a subset of prompts in case there's a lot of imgs / prompts (possibly use img_grids for some gpt4-v calls)
- currently some sub-optimal stuff can happen in preprocess() when there's less than 3 or more than 45 imgs, try to improve this
- Test if timesteps = torch.randint() can be improved: look at sdxl training code! (see https://github.com/huggingface/diffusers/blob/main/examples/advanced_diffusion_training/train_dreambooth_lora_sdxl_advanced.py#L1263, https://arxiv.org/pdf/2206.00364.pdf)
- Fix aspect_ratio bucketing in the dataloader (see https://github.com/kohya-ss/sd-scripts)
- 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
- improve data augmentation, eg by adding outpainted, smaller versions of faces / objects
- 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)
Small, minor tweaks:
- preprocess.py: the imgs are first auto-captioned and then cropped, this is not ideal, swap this around!
Bigger improvements:
- add stronger token regularization (eg CelebBasis spanning basis):
- Add multi-token training
- implement perfusion ideas (key locking with superclass): https://research.nvidia.com/labs/par/Perfusion/
- 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:
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: 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"
-12
View File
@@ -1,12 +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/xander_big.zip" \
-i concept_mode="face" \
-i sd_model_version="sdxl" \
-i max_train_steps="360" \
-i caption_model="blip" \
-i debug="False" \
-i seed="0"
File diff suppressed because it is too large Load Diff
-533
View File
@@ -1,533 +0,0 @@
{
"last_node_id": 12,
"last_link_id": 23,
"nodes": [
{
"id": 7,
"type": "CLIPTextEncode",
"pos": [
413,
389
],
"size": {
"0": 425.27801513671875,
"1": 180.6060791015625
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 16
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
6
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"text, watermark"
]
},
{
"id": 8,
"type": "VAEDecode",
"pos": [
1209,
188
],
"size": {
"0": 210,
"1": 46
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "samples",
"type": "LATENT",
"link": 7
},
{
"name": "vae",
"type": "VAE",
"link": 8
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
9
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "VAEDecode"
}
},
{
"id": 12,
"type": "Reroute",
"pos": [
220,
14
],
"size": [
75,
26
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 22
}
],
"outputs": [
{
"name": "",
"type": "MODEL",
"links": [
19
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 3,
"type": "KSampler",
"pos": [
863,
186
],
"size": {
"0": 315,
"1": 262
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 19
},
{
"name": "positive",
"type": "CONDITIONING",
"link": 4
},
{
"name": "negative",
"type": "CONDITIONING",
"link": 6
},
{
"name": "latent_image",
"type": "LATENT",
"link": 2
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
7
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "KSampler"
},
"widgets_values": [
1,
"fixed",
25,
8,
"euler",
"normal",
1
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
473,
609
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
768,
768,
1
]
},
{
"id": 4,
"type": "CheckpointLoaderSimple",
"pos": [
-467,
120
],
"size": {
"0": 315,
"1": 98
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
10
],
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
12
],
"slot_index": 1
},
{
"name": "VAE",
"type": "VAE",
"links": [
8
],
"slot_index": 2
}
],
"properties": {
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"juggernaut_reborn.safetensors"
]
},
{
"id": 6,
"type": "CLIPTextEncode",
"pos": [
415,
186
],
"size": {
"0": 422.84503173828125,
"1": 164.31304931640625
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 15
}
],
"outputs": [
{
"name": "CONDITIONING",
"type": "CONDITIONING",
"links": [
4
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "CLIPTextEncode"
},
"widgets_values": [
"a photo of embedding:xander_sd15_embedding on the beach "
]
},
{
"id": 11,
"type": "Reroute",
"pos": [
220,
49
],
"size": [
75,
26
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "",
"type": "*",
"link": 23
}
],
"outputs": [
{
"name": "",
"type": "CLIP",
"links": [
15,
16
],
"slot_index": 0
}
],
"properties": {
"showOutputText": false,
"horizontal": false
}
},
{
"id": 9,
"type": "SaveImage",
"pos": [
347,
-270
],
"size": {
"0": 210,
"1": 270
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 9
}
],
"properties": {},
"widgets_values": [
"ComfyUI"
]
},
{
"id": 10,
"type": "LoraLoader",
"pos": [
-72,
-229
],
"size": {
"0": 254.95774841308594,
"1": 127.86701202392578
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 10
},
{
"name": "clip",
"type": "CLIP",
"link": 12
}
],
"outputs": [
{
"name": "MODEL",
"type": "MODEL",
"links": [
22
],
"shape": 3,
"slot_index": 0
},
{
"name": "CLIP",
"type": "CLIP",
"links": [
23
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LoraLoader"
},
"widgets_values": [
"xander_sd15_lora.safetensors",
0.6,
0.6
]
}
],
"links": [
[
2,
5,
0,
3,
3,
"LATENT"
],
[
4,
6,
0,
3,
1,
"CONDITIONING"
],
[
6,
7,
0,
3,
2,
"CONDITIONING"
],
[
7,
3,
0,
8,
0,
"LATENT"
],
[
8,
4,
2,
8,
1,
"VAE"
],
[
9,
8,
0,
9,
0,
"IMAGE"
],
[
10,
4,
0,
10,
0,
"MODEL"
],
[
12,
4,
1,
10,
1,
"CLIP"
],
[
15,
11,
0,
6,
0,
"CLIP"
],
[
16,
11,
0,
7,
0,
"CLIP"
],
[
19,
12,
0,
3,
0,
"MODEL"
],
[
22,
10,
0,
12,
0,
"*"
],
[
23,
10,
1,
11,
0,
"*"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": {
"0": 513.8734070325743,
"1": 351.4824273966635
}
}
},
"version": 0.4
}
+67 -11
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
@@ -382,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):
-656
View File
@@ -1,656 +0,0 @@
import fnmatch
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 typing import Union, Iterable, List, Dict, Tuple, Optional, cast
#from diffusers.training_utils import cast_training_params
from trainer.utils.utils import *
from trainer.checkpoint import save_checkpoint
from trainer.embedding_handler import TokenEmbeddingsHandler
from trainer.dataset import PreprocessedDataset
from trainer.config import TrainingConfig
from trainer.models import print_trainable_parameters, load_models
from trainer.loss import compute_diffusion_loss, compute_grad_norm, ConditioningRegularizer
from trainer.inference import render_images, get_conditioning_signals
from trainer.preprocess import preprocess
from trainer.utils.io import make_validation_img_grid
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)
from trainer.ti_cross_attn_loss import init_daam_loss
pipe, daam_loss = init_daam_loss(
pipeline=pipe
)
config.sd_model_version = sd_model_version
config.pretrained_model["version"] = sd_model_version
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
)
optimizer_unet = get_unet_optimizer(
prodigy_d_coef=config.prodigy_d_coef,
prodigy_growth_factor=config.unet_prodigy_growth_factor,
lora_weight_decay=config.lora_weight_decay,
use_dora=config.use_dora,
unet_trainable_params=unet_trainable_params,
optimizer_name=config.unet_optimizer_type
)
print_trainable_parameters(unet, model_name = 'unet')
for i, text_encoder in enumerate(text_encoders):
if text_encoder is not None:
print_trainable_parameters(text_encoder, model_name = f'text_encoder_{i}')
train_dataset = PreprocessedDataset(
input_dir,
pipe,
vae.float(),
size = config.train_img_size,
do_cache=config.do_cache,
substitute_caption_map=config.token_dict,
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
train_batch_size=config.train_batch_size
)
# offload the vae to cpu:
vae = vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
print(f"# Trainer : Loaded dataset, do_cache: {config.do_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=config.train_batch_size,
shuffle=True,
num_workers=config.dataloader_num_workers,
)
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps)
num_update_steps_per_epoch = math.ceil(len(train_dataloader))
if config.max_train_steps is None:
config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch
config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch)
total_batch_size = config.train_batch_size * config.gradient_accumulation_steps
print(f"--- Num samples = {len(train_dataset)}")
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': []}
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 value of cold (pre-warmup) optimizer lr:
if config.sd_model_version == "sdxl":
if config.is_lora: # let textual_inversion do the work first!
base_lr = 1.0e-5
else:
base_lr = 3.0e-5
elif config.sd_model_version == "sd15":
# let lora training kick in soonish (pure ti for sd15 is not working super well in my tests)
base_lr = 1.0e-4
#######################################################################################################
"""
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
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": # Update ti_learning rate gradually:
if optimizers['textual_inversion'] is not None:
optimizers['textual_inversion'].param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0
# warmup the ti-lr:
if config.ti_lr_warmup_steps > 0:
warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0)
optimizers['textual_inversion'].param_groups[0]['lr'] *= warmup_f
if config.freeze_ti_after_completion_f <= completion_f:
optimizers['textual_inversion'].param_groups[0]['lr'] *= 0
if optimizers['text_encoders'] is not None:
optimizers['text_encoders'].param_groups[0]['lr'] = config.text_encoder_lora_lr * (1 - completion_f) ** 2.0
# 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_lr) ** (global_step / config.unet_lr_warmup_steps)
# Apply the exponential learning rate
optimizers['unet'].param_groups[0]['lr'] = base_lr * exp_factor
if not config.aspect_ratio_bucketing:
captions, vae_latent, mask = batch
else:
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
captions = list(captions)
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]
"""
distirbution shift loss
"""
non_ti_heatmaps = []
ti_heatmaps = []
ti_token_indices = [0,1]
batch_index = 0
token_strings = [
pipe.tokenizer.decode(x)
for x in pipe.tokenizer.encode(captions[batch_index])
]
for text_token_index in range(1, len(token_strings)-1):
if text_token_index in ti_token_indices:
ti_heatmaps.append(
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
)
else:
# we unsqueeze because we'll stack them together and then calculate the min, max and the mean
non_ti_heatmaps.append(
daam_loss.get_the_daam_heatmap(text_token_index = text_token_index).unsqueeze(0)
)
non_ti_heatmaps = torch.cat(
non_ti_heatmaps,
dim = 0
)
ti_heatmaps = torch.cat(
ti_heatmaps,
dim = 0
)
non_ti_dist = {
"mean": non_ti_heatmaps.mean(),
"min": non_ti_heatmaps.min(),
"max": non_ti_heatmaps.min()
}
ti_dist = {
"mean": ti_heatmaps.mean(),
"min": ti_heatmaps.min(),
"max": ti_heatmaps.min()
}
dist_loss = (non_ti_dist["mean"] - ti_dist["mean"].to(non_ti_dist["mean"].device)) ** 2
if global_step % 20 == 0:
batch_index = 0
folder = "./heatmaps"
fig = plt.figure()
token_strings = [
pipe.tokenizer.decode(x)
for x in pipe.tokenizer.encode(captions[batch_index])
]
plot_token_indices = range(len(token_strings))
fig, ax = plt.subplots(nrows=1, ncols=len(plot_token_indices), figsize = (int(3 * len(plot_token_indices)) , 10))
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
im = ax[idx].imshow(heatmap)
ax[idx].set_title(f"{token_strings[text_token_index]}\n timestep: {timesteps[batch_index].item()}\nmax: {heatmap.max().item()}\nmin: {heatmap.min().item()}\nnorm: {heatmap.norm().item()}")
ax[idx].axis("off")
fig.savefig(
os.path.join(
folder,
f"{global_step}.jpg"
)
)
plt.close(fig)
"""
histogram to visualize the distributions of the cross attention values for each text token on the image space
"""
fig = plt.figure()
fig.suptitle(f"Dist loss: {dist_loss.item()}")
plot_token_indices = range(1, len(token_strings)-1)
for idx, text_token_index in enumerate(plot_token_indices):
heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index)[batch_index].cpu().detach().float()
plt.hist(heatmap.reshape(-1), bins = 30, label = token_strings[text_token_index], alpha = 0.5)
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
plt.xlabel("Value")
plt.ylabel("Number of instances")
plt.grid()
# Adjust the layout to prevent the legend from being cut off
plt.tight_layout()
fig.savefig(
os.path.join(
folder,
f"{global_step}_heatmap.jpg"
),
bbox_inches='tight' # This ensures the legend is not cut off when saving
)
plt.close(fig) # Close the figure to free up memory
# Compute the loss:
loss = compute_diffusion_loss(config, model_pred, noise, noisy_latent, mask, noise_scheduler, timesteps)
losses['img_loss'].append(loss.item())
if 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 + 1e-4 * dist_loss
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())
# Print some statistics:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0:
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:
token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings()
for idx, text_encoder in enumerate(text_encoders):
if text_encoder is None:
continue
n = len(token_embeddings[f'txt_encoder_{idx}'])
for i in range(n):
token = trainable_tokens[f'txt_encoder_{idx}'][i]
# Strip any backslashes from the token name:
token = token.replace("/", "_")
embedding = token_embeddings[f'txt_encoder_{idx}'][i]
plot_torch_hist(embedding, global_step, os.path.join(config.output_dir, 'ti_embeddings') , f"enc_{idx}_tokid_{i}: {token}", min_val=-0.05, max_val=0.05, ymax_f = 0.05, color = 'red')
embedding_handler.print_token_info()
if config.is_lora: # plotting this hist for full unet parameters can run OOM
plot_torch_hist(unet_lora_parameters, global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
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,
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//50) == 0:
progress = (global_step / config.max_train_steps) + 0.05
#print_system_info()
print(f"\n---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r", flush = True)
yield np.min((progress, 1.0))
if global_step > config.max_train_steps:
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,
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 :)")
-130
View File
@@ -1,130 +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": "."}),
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"lora_name": ("STRING", {"default": "Eden_LoRa"}),
"mode": (["style", "face", "object"], ),
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
"ti_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"unet_lr": ("FLOAT", {"default": 0.001, "min": 0.0001, "max": 0.01, "step": 0.0001}),
"lora_rank": ("INT", {"default": 16, "min": 1, "max": 64}),
"use_dora": ("BOOLEAN", {"default": False}),
"n_tokens": ("INT", {"default": 2, "min": 1, "max": 3}),
"debug_mode": ("BOOLEAN", {"default": False}),
"checkpointing_steps": ("INT", {"default": 200, "min": 10, "max": 2000}),
"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 = "eden_lora",
mode = "style",
seed = 0,
resolution = 521,
train_batch_size = 4,
max_train_steps = 400,
ti_lr = 0.001,
unet_lr = 0.001,
lora_rank = 16,
use_dora = False,
n_tokens = 2,
debug_mode = False,
checkpointing_steps = 1000,
):
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("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,
lora_training_urls=training_images_folder_path,
concept_mode=mode,
ckpt_path=ckpt_path,
seed=seed,
resolution=resolution,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=checkpointing_steps,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
use_dora=use_dora,
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
debug=debug_mode,
)
pbar = comfy.utils.ProgressBar(100)
with torch.inference_mode(False):
train_generator = train(config=config)
while True:
try:
progress_f = next(train_generator)
pbar.update_absolute(progress_f * 100)
except StopIteration as e:
config, output_save_dir = e.value # Capture the return value
break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
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 image:
grid_image = Image.open(validation_grid_img_path)
grid_image = np.array(grid_image).astype(np.float32) / 255.0
grid_image = torch.from_numpy(grid_image)[None,]
final_msg = f"LoRa trained in {config.job_time/60:.1f} minutes. Files saved at {output_save_dir}"
return (grid_image, lora_path, embedding_path, final_msg)
+353 -72
View File
@@ -7,20 +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.models import pretrained_models
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()
@@ -29,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
@@ -52,93 +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=400
),
resolution: int = Input(
description="Square pixel resolution which your images will be resized to for training, highly recommended: 512 or 640",
default=512
),
train_batch_size: int = Input(
description="Batch size (per device) for training (dont increase unless running on a BIG GPU)",
default=4
),
unet_lr: float = Input(
description="final learning rate of unet (after warmup), increasing this usually leads to strong overfitting",
default=0.001
),
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
),
use_dora: bool = Input(
description="Use Dora instead of LoRa",
default=False,
),
n_tokens: int = Input(
description="How many new tokens to train (highly recommended to leave this at 2)",
ge=1, le=3, default=2
),
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,
train_batch_size=train_batch_size,
num_train_epochs=num_train_epochs,
max_train_steps=max_train_steps,
checkpointing_steps=10000,
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,
use_dora=use_dora,
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:
@@ -146,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)
@@ -162,18 +443,18 @@ class Predictor(BasePredictor):
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
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)
+198 -401
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 = 50
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,68 +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."},
@@ -302,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):
text = text.strip()
gpt_cleanup_worked = False
gpt_concept_description = None
if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client:
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...")
@@ -374,14 +395,16 @@ def post_process_captions(captions, text, concept_mode, job_seed):
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(
@@ -393,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)):
@@ -419,10 +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)
del model
gc.collect()
torch.cuda.empty_cache()
return captions
def encode_image(image_path):
@@ -440,51 +454,6 @@ def prep_img_for_gpt_api(pil_img, max_size=(512, 512)):
os.remove(output_path)
return base64_image
def gpt4_v_get_description(config, images):
if config.concept_mode == "object":
description = "object"
prompt = "Give a concise visual descriptioni of the object/figure/thing that all the grid-images have in common with at most 10 words. Dont start with statements like 'The image features...', just describe what you see."
elif config.concept_mode == "face":
description = "face"
prompt = "All the grid images depict a single person. Visually describe this person with at most 10 words. Dont start with statements like 'The image features...', just describe what you see. (eg an asian woman with long black hair)"
elif config.concept_mode == "style":
description = ""
prompt = "All these images share a common aesthetic style. Describe this style with at most 7 words. Dont start with statements like 'The image features...', just describe what you see. (eg impressionism collage surrealism)"
if not OPENAI_API_KEY:
print(f"Skipping GPT-4 Vision description because OPENAI_API_KEY is not set.")
return description
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {OPENAI_API_KEY}"
}
# TODO sample a grid img:
# .... TODO
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
payload = {
"model": "gpt-4o",
"messages": [
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{base64_image}", "detail": "high"}}
]
}
],
"max_tokens": 60
}
response = requests.post("https://api.openai.com/v1/chat/completions", headers=headers, json=payload)
answer = response.json()["choices"][0]["message"]["content"]
return captions
def gpt4_v_caption_dataset(
images, captions,
batch_size=4,
@@ -494,7 +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 = "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",
@@ -505,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",
@@ -515,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}
@@ -557,6 +519,49 @@ def caption_dataset(
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)
@@ -635,53 +640,9 @@ def augment_image(image):
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",
@@ -691,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,
@@ -706,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)
@@ -714,7 +677,7 @@ 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)
@@ -730,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.")
@@ -761,47 +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) < 40:
print(f"Adding LR flips... (doubling the number of images from {n_training_imgs} to {n_training_imgs*2})")
images = images + [image.transpose(Image.FLIP_LEFT_RIGHT) for image in images]
captions = captions + captions
# It's nice if we can achieve the gpt pass, so pre-augment the images if there's very few:
# Ensure we have at least 'augment_imgs_up_to_n' images through augmentation
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < MIN_GPT_PROMPTS:
print(f"Adding augmented version of each training img...")
aug_imgs.extend([augment_image(image) for image in images])
aug_caps.extend(captions)
images.extend(aug_imgs)
captions.extend(aug_caps)
# It's nice if we can achieve the gpt pass, so if we're not losing too much, cut-off the n_images to just match what we're allowed to give to gpt:
if (len(images) > MAX_GPT_PROMPTS) and (len(images) < MAX_GPT_PROMPTS*1.33):
images = images[:MAX_GPT_PROMPTS-1]
captions = captions[:MAX_GPT_PROMPTS-1]
if len(images) > 50 and caption_model != "blip":
print(f"Captioning a lot of ({len(images)}) images --> falling back to using blip!")
caption_model = "blip"
# 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)
captions, trigger_text, gpt_concept_name = post_process_captions(captions, caption_text, concept_mode, seed)
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
while len(images) + len(aug_imgs) < augment_imgs_up_to_n:
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)
@@ -809,30 +728,25 @@ def load_and_save_masks_and_captions(
images.extend(aug_imgs)
captions.extend(aug_caps)
if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")):
print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}")
mask_target_prompts = gpt_concept_description
if (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 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
@@ -842,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):
@@ -874,22 +780,9 @@ 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]
print("Final 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)
@@ -913,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.
@@ -1058,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]]
@@ -1133,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
-22
View File
@@ -1,22 +0,0 @@
torch==2.1.0
torchaudio==2.1.0
torchvision==0.16.0
transformers==4.38.0
diffusers==0.26.0
tokenizers==0.15.2
huggingface-hub==0.22.2
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
-161
View File
@@ -1,161 +0,0 @@
"""
Faces:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
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://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
"""
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 = "beeple"
caption_prefix = ""
mask_target_prompts = ""
n_exp = 200 # how many random experiment settings to generate
min_hamming_distance = 1 # min_n_params that have to be different from any previous experiment to be scheduled
nohup = True
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": ["sdxl"],
"lora_training_urls": [
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.8],
"disable_ti": ['false', 'true'],
"seed": [0],
"resolution": [512],
"train_batch_size": [4],
"n_sample_imgs": [8],
"max_train_steps": [1200],
"checkpointing_steps": [200],
"gradient_accumulation_steps": [1],
"n_tokens": [2],
"ti_lr": [0.001],
"ti_weight_decay": [0.001],
"l1_penalty": [0.0],
"token_warmup_steps": [0],
"tok_cov_reg_w": [2000],
"unet_lr": [0.0002, 0.00005],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16],
"use_dora": ['false'],
"unet_optimizer_type": ['AdamW8bit'],
"is_lora": ['false'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["blip", "gpt4-v"],
"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 = 200
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)
# Make some final adjustments to the experiment settings before saving to disk:
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
with open(config_filename, "w") as f:
json.dump(experiment_settings, f, indent=4)
break
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}")
-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
"""
-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"
-8
View File
@@ -1,8 +0,0 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=0"
python main.py train_configs/training_args_face_sdxl.json
python main.py train_configs/training_args_face_sd15.json
python main.py train_configs/training_args_object.json
python main.py train_configs/training_args_style_sd15.json
python main.py train_configs/training_args_style_sdxl.json
-27
View File
@@ -1,27 +0,0 @@
{
"name": "xander_test",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 4,
"max_train_steps": 200,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
-28
View File
@@ -1,28 +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_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"n_tokens": 2,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.00,
"lora_rank": 4,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -1,27 +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_5.zip",
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 300,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -1,27 +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_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"disable_ti": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
-28
View File
@@ -1,28 +0,0 @@
{
"name": "banny_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 800,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -1,27 +0,0 @@
{
"name": "clipx_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
@@ -1,28 +0,0 @@
{
"name": "clipx_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+2
View File
@@ -0,0 +1,2 @@
from .trainer import Trainer
from .config import TrainerConfig
-276
View File
@@ -1,276 +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
def remove_delimiter_characters(name: str):
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
# Convert to WebUI format
def convert_pytorch_lora_safetensors_to_webui(
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
+48 -169
View File
@@ -1,182 +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
import torch
class ModelPaths:
def __init__(self):
self.paths = {
"BLIP": "./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 download urls in case no local model is found:
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-XL-v6/resolve/main/juggernautXL_version6Rundiffusion.safetensors"
SDXL_URL = "https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_base_1.0_0.9vae.safetensors"
SD15_URL = "https://huggingface.co/KamCastle/jugg/resolve/main/juggernaut_reborn.safetensors"
pretrained_models = {
"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
caption_model: Literal["gpt4-v", "blip"] = "blip"
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
num_train_epochs: int = 10000
max_train_steps: int = 360
checkpointing_steps: int = 10000
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 1.0e-3
prodigy_d_coef: float = 1.0
unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
lora_weight_decay: float = 0.002
ti_lr: float = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0
ti_optimizer: Literal["adamw", "prodigy"] = "adamw"
freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done
cond_reg_w: float = 0.0e-5
tok_cond_reg_w: float = 0.0e-5
tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens
off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std
l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse
noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images
snr_gamma: float = 5.0
lora_alpha_multiplier: float = 1.0
lora_rank: int = 12
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
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 2
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
token_dict: dict = {"TOK": "<s0><s1>"}
device: str = "cuda:0"
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
do_cache: bool = True
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
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
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
snr_gamma: float = 5.0
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}
# add some metrics to the foldername:
lora_str = "dora" if self.use_dora else "lora"
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
if not self.name:
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
self.output_dir = self.output_dir + f"/{self.name}/" + f"{timestamp_short}-{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
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.concept_mode == "face":
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
self.left_right_flip_augmentation = False # always disable lr flips for face mode!
self.mask_target_prompts = "face"
#self.use_face_detection_instead = True
if not self.sample_imgs_lora_scale:
if self.sd_model_version == "sdxl":
self.sample_imgs_lora_scale = 0.7
else:
self.sample_imgs_lora_scale = 0.85
if self.use_dora:
print(f"Disabling L1 penalty and LoRA weight decay for DORA training.")
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 -168
View File
@@ -1,191 +1,171 @@
import os
import torch
import numpy as np
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,
do_cache: bool = False,
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)
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.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 do_cache:
print("Caching latents, masks and captions...\n")
self.vae_latents = []
self.masks = []
self.do_cache = True
for idx in range(len(self.data)):
if len(self.data) < 25:
print(self.captions[idx])
vae_latent, mask = self._process(idx)
self.vae_latents.append(vae_latent)
self.masks.append(mask)
print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.")
del self.vae_encoder
else:
self.do_cache = False
if aspect_ratio_bucketing:
print("Using aspect ratio bucketing.")
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()
print(f"Got bucket batch: {indices}, resolution: {resolution}")
tok1, tok2, vae_latents, masks = [], [], [], []
for idx in indices:
if self.tokenizer_2 is None:
t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution)
else:
(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, device=self.vae_encoder.device
)
else:
image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to(
dtype=self.vae_encoder.dtype, device=self.vae_encoder.device
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).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, device=self.vae_encoder.device
)
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, device=self.vae_encoder.device
)
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()
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(), self.masks[idx]
else: # This code pathway has not been tested in a long time and might be broken
caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution)
vae_latent = vae_latent.sample() * self.vae_scaling_factor
return caption, vae_latent.squeeze(), mask
# 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)
-497
View File
@@ -1,497 +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 = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
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()
self.fix_embedding_std(config.off_ratio_power)
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 fix_embedding_std(self, off_ratio_power=0.1):
if off_ratio_power == 0.0:
return
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
if text_encoder is None:
idx += 1
continue
# Get the standard deviation target and current embeddings.
target_std = self.embeddings_settings[f"std_token_embedding_{idx}"]
embeddings, _ = self.get_trainable_embeddings()
new_embeddings = embeddings[f'txt_encoder_{idx}']
assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!"
new_stds = new_embeddings.std(dim=1)
#off_ratios = target_std.float() / new_stds.float()
off_ratios = target_std / new_stds
# Check if off_ratios are within an acceptable range.
if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1):
# Convert the pytorch tensor into a list of python floats:
off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3)
print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...")
# Adjust embeddings using the computed ratios.
index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"]
index_updates = ~index_no_updates
multiplier_values = off_ratios**off_ratio_power
multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings)
text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values
idx += 1
def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder):
# Assuming new tokens are of the format <s_i>
self.inserting_toks = [f"<s{i}>" for i in range(loaded_embeddings.shape[0])]
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)
-489
View File
@@ -1,489 +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 = (
config.crops_coords_top_left_h,
config.crops_coords_top_left_w,
)
if pipe.text_encoder_2 is None:
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
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:
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,
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 (
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,
)
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,
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,
)
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
-375
View File
@@ -1,375 +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
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.003, pipe=None):
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,2:-2,:]) * noise_sigma
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]
"""
-97
View File
@@ -1,97 +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, keep_vae_float32 = False):
# check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
try:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sdxl"
except:
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)
if keep_vae_float32:
vae.to(device, dtype=torch.float32)
else:
vae.to(device, dtype=weight_dtype)
if weight_dtype != torch.float32:
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but may not be ideal for training..?")
unet.to(device, dtype=weight_dtype)
text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype)
tokenizer_two = text_encoder_two = None
if sd_model_version == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
text_encoder_two.to(device, dtype=weight_dtype)
return (
pipe,
tokenizer_one,
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()
line_delimiter = "#" * 80
print(line_delimiter)
print(
f"Trainable {model_name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
)
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]))
-289
View File
@@ -1,289 +0,0 @@
from functools import reduce
from diffusers import StableDiffusionXLPipeline
from diffusers.models.attention_processor import AttnProcessor2_0, Attention
from typing import Optional
import torch
import torch.nn as nn
from diffusers.utils.deprecation_utils import deprecate
import torch.nn.functional as F
import math
from einops.layers.torch import Reduce
from torchtyping import TensorType
from einops import rearrange
# 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 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 compute_single_token_loss(self, text_token_index: list[int], reduce = False):
loss = {}
cross_attention_scores = self.get_all_cross_attention_scores()
for name, cross_attention_map in cross_attention_scores.items():
"""
cross_attention_map.shape: (batch, image_patches, text_tokens)
"""
assert cross_attention_map.ndim == 3
loss[name] = cross_attention_map[:,:,text_token_index].norm() / cross_attention_map.shape[1]
if reduce:
all_losses = list(loss.values())
return sum(all_losses)/len(all_losses)
else:
return loss
def compute_loss(self, text_token_indices: list[int], reduce = False):
losses = []
for text_token_index in text_token_indices:
losses.append(
self.compute_single_token_loss(
text_token_index=text_token_index,
reduce = True
)
)
if reduce:
return sum(losses)/len(losses)
else:
return losses
def get_image_heatmap(self, text_token_index: int, layer_name: str) -> TensorType["batch", "height", "width"]:
cross_attention_scores = self.get_all_cross_attention_scores()
assert layer_name in list(cross_attention_scores.keys())
cross_attention_scores_single_token = cross_attention_scores[layer_name][:,:,text_token_index]
assert cross_attention_scores_single_token.ndim == 2 ## batch, hw
heatmap = rearrange(
cross_attention_scores_single_token,
"batch (height width) -> batch height width",
height = int(math.sqrt(cross_attention_scores_single_token.shape[1])),
width = int(math.sqrt(cross_attention_scores_single_token.shape[1]))
)
return heatmap
def get_the_daam_heatmap(self, text_token_index: int) ->TensorType["batch", "height", "width"]:
all_heatmaps = []
for layer_name in self.layer_names:
heatmap = self.get_image_heatmap(
text_token_index=text_token_index,
layer_name=layer_name
)
all_heatmaps.append(heatmap)
## each heatmap has a shape: batch, h, w where h=w
## now find the maximum possible height and width across all heatmaps
max_height = max(heatmap.shape[1] for heatmap in all_heatmaps)
max_width = max(heatmap.shape[2] for heatmap in all_heatmaps)
## now resize all_heatmaps to (batch, max_height, max_width) using F.interpolate
resized_heatmaps = [
F.interpolate(input = x.unsqueeze(1), size = (max_height, max_width)).squeeze(1)
for x in all_heatmaps
]
return sum(resized_heatmaps)
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: StableDiffusionXLPipeline)-> tuple[StableDiffusionXLPipeline, DAAMLoss]:
assert isinstance(pipeline, StableDiffusionXLPipeline)
## 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:
# print(f"Replacing: {name}")
# 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'}
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()
@@ -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"