45 Commits
Author SHA1 Message Date
mayukhdeb fb8ee98a1c penalize mean 2024-08-02 04:31:01 -07:00
mayukhdeb 7ff02a35ce smaller daam loss scale 2024-08-02 04:13:23 -07:00
mayukhdeb 21200b41b9 naive distribution loss, did not work :( 2024-08-02 03:17:32 -07:00
mayukhdeb c21605bb95 also plot token as string on title 2024-08-01 00:49:47 -07:00
mayukhdeb 4880e3f096 vis token wise daam maps 2024-08-01 00:39:49 -07:00
mayukhdeb 35e76fdef0 wip reproduce daam 2024-07-30 23:01:14 -07:00
mayukhdeb 871e425f3e obtain single layer heatmap 2024-07-30 22:50:56 -07:00
mayukhdeb e6534f8d1d baby steps 2024-07-30 22:50:22 -07:00
mayukhdeb b5ad0b659e watch daam loss and plot norms into heatmaps dir 2024-07-29 10:46:11 -07:00
mayukhdeb d6a205cc7b ahijack AttnProcessor2_0 2024-07-29 10:12:43 -07:00
mayukhdeb f21efd7840 train TI on 2 tokens 2024-07-29 10:11:15 -07:00
mayukhdeb 29855d94ba ignore notebook checkpoints folder 2024-07-22 12:03:25 -07:00
xander 006a3b9750 add ti config 2024-07-22 20:53:49 +02:00
aiXander b166f7c274 merge 2024-07-22 11:47:38 -07:00
aiXander f6cfd6d0bc fix typo 2024-07-22 11:46:33 -07:00
xander 6a4fa1ef8d add progressbar to node 2024-07-18 21:04:43 +02:00
xander 3cd17086f3 trainer comfyui v1 2024-07-18 05:50:04 +02:00
xander 75137a7539 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:07:07 +02:00
xander ee91f8c0c7 update naming conventions 2024-07-18 04:07:00 +02:00
xander 598c4204be add test config 2024-07-18 04:05:10 +02:00
xander 8da6c91af8 Merge branch 'main' of https://github.com/edenartlab/sd-lora-trainer into main 2024-07-18 04:00:30 +02:00
xander 01523350e1 cleanup save dir 2024-07-18 04:00:21 +02:00
xander c891203365 update gitignore 2024-07-18 04:00:08 +02:00
xander 85599dc725 minor changes 2024-07-18 03:43:43 +02:00
xander af16d1edcc update configs 2024-07-18 03:26:51 +02:00
xander b3db4ff3bb Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-07-18 03:16:43 +02:00
xander 2a0761ae75 use gpt-4o 2024-07-18 03:16:41 +02:00
xander d8ee9ee522 use gpt-4o 2024-07-18 03:14:31 +02:00
xander 4f084d251c updates for comfyui node 2024-07-18 03:12:11 +02:00
aiXander fb521e9dbd add print flush 2024-07-16 04:39:29 -07:00
aiXander f937106f6c push changes 2024-07-16 04:04:34 -07:00
Gene Kogan 5f8fe5c4c8 Update requirements.txt 2024-07-16 03:44:59 -07:00
Gene Kogan 2f5eaeba7a Update cog.yaml 2024-07-16 03:44:35 -07:00
aiXander 7d8c845765 update yaml and print deps in config 2024-07-12 05:16:06 -07:00
aiXander 340336b53d print config pre training start 2024-07-11 06:11:33 -07:00
xander 5abed1487f setup lora_scale automation for validation grid 2024-07-11 13:36:11 +02:00
xander b5cf857b1f push small changes 2024-07-09 18:33:52 +02:00
xander e1813c75d1 updates to avoid OOM when plotting hist 2024-07-06 22:18:47 +02:00
xander 4fde1b7dd6 small tweaks 2024-07-06 18:45:38 +02:00
xander 892cf61c52 add disable_ti flag 2024-07-06 18:43:03 +02:00
xander fdd531d3c5 fix full finetuning and add 8bitadam 2024-07-06 17:42:49 +02:00
xander ef6635fafb update train cmd 2024-07-05 17:11:06 +02:00
xander faf864ce15 tiny tweak to learning rates for SDXL 2024-07-05 17:07:33 +02:00
xander b02221b23c cleanup before sd3 integration 2024-06-13 14:24:32 +02:00
xander d8b3175b57 cleanup before sd3 integration 2024-06-13 13:57:20 +02:00
23 changed files with 886 additions and 283 deletions
+2 -2
View File
@@ -1,16 +1,16 @@
cache
__pycache__
.ipynb_checkpoints/
models
lora_models*
eden_lora_training_runs/
datasets
*.tar
.env
.cog
.huggingface
train.py
rendered_images*
gridsearch*
+4 -3
View File
@@ -29,7 +29,7 @@ Install all dependencies using
then you can simply run:
`python main.py -c training_args.json`
`python main.py train_configs/training_args.json`
to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job.
@@ -44,8 +44,9 @@ 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 `sudo cog build`
3. Run a training run with `sudo sh cog_test_train.sh`
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
+4 -7
View File
@@ -3,17 +3,14 @@
build:
gpu: true
cuda: "11.8"
python_version: "3.9"
cuda: "12.1"
python_version: "3.11"
system_packages:
- "ffmpeg"
- "libgl1-mesa-glx"
- "libegl1-mesa-dev"
- "libsm6"
- "libxext6"
python_requirements: requirements.txt
run:
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
predict: "predict.py:Predictor"
image: "r8.im/edenartlab/sdxl-lora-trainer"
+1 -1
View File
@@ -1,5 +1,5 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=2"
GPU_ID="device=3"
cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \
+178 -69
View File
@@ -12,7 +12,6 @@ import torch
import torch.utils.checkpoint
from tqdm import tqdm
import prodigyopt
from typing import Union, Iterable, List, Dict, Tuple, Optional, cast
#from diffusers.training_utils import cast_training_params
@@ -26,6 +25,7 @@ from trainer.loss import compute_diffusion_loss, compute_grad_norm, Conditioning
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,
@@ -34,10 +34,28 @@ from trainer.optimizer import (
get_unet_optimizer
)
def train(
config: TrainingConfig,
):
def train(config: TrainingConfig):
seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
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,
@@ -58,19 +76,6 @@ def train(
if config.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(config.pretrained_model, config.device, weight_dtype, keep_vae_float32=0)
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
text_encoders = [text_encoder_one, text_encoder_two],
@@ -113,22 +118,26 @@ def train(
embedding_handler.make_embeddings_trainable()
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
if not config.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
@@ -194,7 +203,7 @@ def train(
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")
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
global_step = 0
last_save_step = 0
@@ -216,10 +225,12 @@ def train(
# default value of cold (pre-warmup) optimizer lr:
if config.sd_model_version == "sdxl":
# let textual_inversion do the work first!
base_lr = 0.5e-5
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
# let lora training kick in soonish (pure ti for sd15 is not working super well in my tests)
base_lr = 1.0e-4
#######################################################################################################
@@ -309,7 +320,109 @@ def train(
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())
@@ -320,7 +433,7 @@ def train(
loss += 0.0 * concept_description_loss
losses['concept_description_loss'].append(concept_description_loss.item())
if config.l1_penalty > 0.0:
if config.l1_penalty > 0.0 and unet_lora_parameters:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
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
@@ -329,6 +442,7 @@ def train(
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()
@@ -349,12 +463,6 @@ def train(
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
optimizer_collection.step()
# after every optimizer step, we do some manual intervention of the embeddings to regularize them:
if optimizer_collection.get_lr('textual_inversion') > 0.0:
#embedding_handler.fix_embedding_std(config.off_ratio_power)
pass
optimizer_collection.zero_grad()
#############################################################################################################
@@ -369,7 +477,7 @@ def train(
token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item())
# Print some statistics:
if config.debug and (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > -1:
if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
@@ -390,27 +498,29 @@ def train(
)
last_save_step = global_step
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')
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()
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
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,
@@ -433,14 +543,14 @@ def train(
images_done += config.train_batch_size
global_step += 1
if global_step % (config.max_train_steps//20) == 0:
if global_step % (config.max_train_steps//50) == 0:
progress = (global_step / config.max_train_steps) + 0.05
print_system_info()
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
#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!")
print("Reached max steps, stopping training!", flush = True)
break
# final_save
@@ -471,8 +581,7 @@ def train(
pretrained_model_version=config.pretrained_model["version"]
)
print("Running final inference round...")
if config.debug:
if config.debug and 0:
# Reload the entire pipe from disk + LoRa:
pipe_to_use = None
checkpoint_folder = output_save_dir
@@ -511,13 +620,6 @@ def train(
img_grid_path = make_validation_img_grid(output_save_dir)
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
# Remove unneeded checkpoints if they exist in the output directory:
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_save_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
else:
print(f"Skipping final save, {output_save_dir} already exists")
@@ -531,6 +633,8 @@ def train(
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
@@ -541,6 +645,11 @@ if __name__ == "__main__":
args = parser.parse_args()
config = TrainingConfig.from_json(file_path=args.config_filename)
print("Starting new LoRa training run with config:")
print(config)
print("------------------------------------------")
for progress in train(config=config):
print(f"Progress: {(100*progress):.2f}%", end="\r")
+58 -43
View File
@@ -1,22 +1,17 @@
import os
import shutil
import tarfile
import json
import time
import random
import torch
import numpy as np
import pandas as pd
from PIL import Image
from dotenv import load_dotenv
from main import train
from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig
from trainer.config import TrainingConfig, model_paths
from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything
import folder_paths
import comfy.utils
class Eden_LoRa_trainer:
@classmethod
@@ -24,9 +19,9 @@ class Eden_LoRa_trainer:
return {
"required": {
"training_images_folder_path": ("STRING", {"default": "."}),
"lora_name": ("STRING", {"default": ""}),
"sd_model_version": (["sdxl", "sd15"], ),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
"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}),
@@ -35,17 +30,22 @@ class Eden_LoRa_trainer:
"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 = ("STRING",)
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,
name = lora_name,
concept_mode = "style",
sd_model_version = "sdxl",
def train_lora(self,
training_images_folder_path,
ckpt_name,
lora_name = "eden_lora",
mode = "style",
seed = 0,
resolution = 521,
train_batch_size = 4,
@@ -54,21 +54,31 @@ class Eden_LoRa_trainer:
unet_lr = 0.001,
lora_rank = 16,
use_dora = False,
n_tokens = 2
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="test",
name=lora_name,
lora_training_urls=training_images_folder_path,
concept_mode=concept_mode,
sd_model_version=sd_model_version,
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=10000,
checkpointing_steps=checkpointing_steps,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
@@ -76,40 +86,45 @@ class Eden_LoRa_trainer:
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
debug=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")
out_path = f"{clean_filename(lora_name)}_eden_concept_lora_{int(time.time())}.tar"
directory = cogPath(output_save_dir)
with tarfile.open(out_path, "w") as tar:
print("Adding files to tar...")
for file_path in directory.rglob("*"):
print(file_path)
arcname = file_path.relative_to(directory)
tar.add(file_path, arcname=arcname)
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
print(f"LORA training finished in {config.job_time:.1f} seconds")
print(f"Returning {out_path}")
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")]
return (out_path,)
# find the index of the path containing "_embeddings.safetensors":
for i, path in enumerate(paths):
if "_embeddings.safetensors" in path:
embedding_path = path
else:
lora_path = path
# Load the grid 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)
+21 -16
View File
@@ -1,17 +1,22 @@
torch>=2.1.0
torchvision>=0.16.0
transformers>=4.38.1
diffusers>=0.27.2
ujson>=5.9.0
scipy>=1.12.0
peft>=0.10.0
invisible-watermark>=0.2.0
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.1.0.25
mediapipe>=0.10.11
openai>=1.14.0
python-dotenv
prodigyopt
omegaconf
ujson
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
+28 -20
View File
@@ -35,12 +35,12 @@ def hamming_distance(dict1, dict2):
#######################################################################################
# Setup the base experiment config:
exp_name = "grimes"
exp_name = "beeple"
caption_prefix = ""
mask_target_prompts = ""
n_exp = 200 # how many random experiment settings to generate
min_hamming_distance = 3 # min_n_params that have to be different from any previous experiment to be scheduled
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
@@ -48,43 +48,46 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sd15", "sdxl"],
"sd_model_version": ["sdxl"],
"lora_training_urls": [
"/home/rednax/Documents/datasets/grimes"
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
],
"concept_mode": ['face'],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.8],
"disable_ti": ['false', 'true'],
"seed": [0],
"resolution": [512],
"train_batch_size": [4],
"n_sample_imgs": [6],
"max_train_steps": [400,800],
"checkpointing_steps": [100],
"n_sample_imgs": [8],
"max_train_steps": [1200],
"checkpointing_steps": [200],
"gradient_accumulation_steps": [1],
"n_tokens": [2],
"ti_lr": [0.001,0.0005],
"ti_weight_decay": [0.001,0.0],
"ti_lr": [0.001],
"ti_weight_decay": [0.001],
"l1_penalty": [0.0],
"token_warmup_steps": [0,60],
"token_warmup_steps": [0],
"tok_cov_reg_w": [2000],
"cond_reg_w": [0.01e-5],
"tok_cond_reg_w": [0.01e-5],
"unet_prodigy_growth_factor": [1.05],
"unet_lr": [0.001],
"unet_lr": [0.0002, 0.00005],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16,32],
"use_dora": ['false', 'true'],
"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": [20,40],
"augment_imgs_up_to_n": [40],
"verbose": ['true'],
"debug": ['true']
}
@@ -146,7 +149,12 @@ def generate_sh_script(folder_path, output_sh_path):
# Write a command for each JSON file
for json_file in json_files:
command = f"python main.py {os.path.join(folder_path, json_file)}\n"
file_path = os.path.join("scripts/", folder_path, json_file)
command = f"python main.py {file_path}\n"
if nohup:
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path)
+8
View File
@@ -0,0 +1,8 @@
# 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
@@ -1,25 +1,27 @@
{
"output_dir": "lora_models/object",
"name": "xander_test",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/DOV/lizzo/full body",
"concept_mode": "object",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
"concept_mode": "face",
"seed": 1,
"resolution": 640,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 4,
"max_train_steps": 420,
"max_train_steps": 200,
"token_warmup_steps": 0,
"checkpointing_steps": 60,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"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,
"lora_rank": 12,
"unet_lr": 0.001,
"lora_rank": 16,
"use_dora": false,
"caption_model": "gpt4-v",
"caption_model": "blip",
"debug": true
}
+28
View File
@@ -0,0 +1,28 @@
{
"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
}
@@ -0,0 +1,27 @@
{
"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
}
@@ -0,0 +1,27 @@
{
"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
}
@@ -1,31 +1,28 @@
{
"output_dir": "lora_models/xander_sd15_final",
"name": "banny_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
"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": 600,
"max_train_steps": 800,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": false,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 16,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": false,
"caption_model": "gpt4-v",
"caption_model": "blip",
"debug": true
}
@@ -1,17 +1,15 @@
{
"output_dir": "lora_models/does_best",
"name": "clipx_sd15",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"seed": 1,
"resolution": 640,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 600,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
@@ -1,29 +1,28 @@
{
"output_dir": "lora_models/Journey",
"name": "clipx_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "/home/rednax/Documents/datasets/journey",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"seed": 0,
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 1000,
"max_train_steps": 400,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"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,
"prodigy_d_coef": 1.0,
"unet_prodigy_growth_factor": 1.05,
"lora_rank": 16,
"use_dora": true,
"caption_model": "gpt4-v",
"use_dora": false,
"caption_model": "blip",
"debug": true
}
+13 -3
View File
@@ -135,7 +135,7 @@ def save_checkpoint(
embedding_handler.save_embeddings(
os.path.join(
output_dir,
f"{name}_embeddings.safetensors"
f"{name}_{pretrained_model_version}_embeddings.safetensors"
)
)
@@ -145,7 +145,7 @@ def save_checkpoint(
output_dir, "special_params.json"
)
)
if is_lora:
assert len(unet_lora_parameters) > 0, f"Expected len(unet_lora_parameters) to be greater than zero if is_lora is True"
@@ -184,11 +184,21 @@ def save_checkpoint(
convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
output_filename=os.path.join(output_dir, f"{name}.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,
+46 -8
View File
@@ -3,15 +3,43 @@ from datetime import datetime
from pydantic import BaseModel
import json, time, os
from typing import Literal
from trainer.models import pretrained_models
from trainer.utils.utils import pick_best_gpu_id
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"}
}
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"]
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
@@ -25,7 +53,7 @@ class TrainingConfig(BaseModel):
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy"] = "adamw"
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 1.0e-3
prodigy_d_coef: float = 1.0
@@ -59,10 +87,10 @@ class TrainingConfig(BaseModel):
clipseg_temperature: float = 0.5 # temperature for the CLIPSeg mask
n_sample_imgs: int = 4
name: str = None
output_dir: str = "lora_models/unnamed"
output_dir: str = "eden_lora_training_runs"
debug: bool = False
allow_tf32: bool = True
remove_ti_token_from_prompts: bool = False
disable_ti: bool = False
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 2
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
@@ -74,7 +102,7 @@ class TrainingConfig(BaseModel):
unet_learning_rate: float = 1.0
lr_num_cycles: int = 1
lr_power: float = 1.0
sample_imgs_lora_scale: float = 0.65 # Default lora scale for sampling the validation images
sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0
training_attributes: dict = {}
aspect_ratio_bucketing: bool = False
@@ -93,7 +121,11 @@ class TrainingConfig(BaseModel):
def __init__(self, **data):
super().__init__(**data)
self.pretrained_model = pretrained_models[self.sd_model_version]
if not self.ckpt_path:
self.pretrained_model = pretrained_models[self.sd_model_version]
else:
self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None}
# add some metrics to the foldername:
lora_str = "dora" if self.use_dora else "lora"
@@ -102,7 +134,7 @@ class TrainingConfig(BaseModel):
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"--{timestamp_short}-{self.sd_model_version}_{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
self.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:
@@ -116,6 +148,12 @@ class TrainingConfig(BaseModel):
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.")
+12 -36
View File
@@ -4,47 +4,23 @@ import subprocess
import torch
from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline
############################################################################################################
SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
#SDXL_MODEL_CACHE = "./models/Juggernaut-X-RunDiffusion-NSFW.safetensors"
#SDXL_URL = "https://huggingface.co/RunDiffusion/Juggernaut-X-v10/resolve/main/Juggernaut-X-RunDiffusion-NSFW.safetensors"
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
#SD15_MODEL_CACHE = "./models/DreamShaper_6.31_BakedVae.safetensors"
#SD15_URL = "https://huggingface.co/Lykon/DreamShaper/resolve/main/DreamShaper_6.31_BakedVae.safetensors"
#SD15_MODEL_CACHE = "./models/photon_v1.safetensors"
#SD15_URL = "https://civitai.com/api/download/models/90072"
pretrained_models = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL, "version": "sd15"}
}
############################################################################################################
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
if not isinstance(pretrained_model, dict) or 'path' not in pretrained_model or 'version' not in pretrained_model:
raise ValueError("pretrained_model must be a dict with 'path' and 'version' keys")
# check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
print(f"Loading model weights from {pretrained_model['path']} with dtype: {weight_dtype}...")
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
if pretrained_model['version'] == "sd15":
pipe = StableDiffusionPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
else:
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)
@@ -60,14 +36,14 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
else:
vae.to(device, dtype=weight_dtype)
if weight_dtype != torch.float32:
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but might not be for training..")
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 pretrained_model['version'] == "sdxl":
if sd_model_version == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
@@ -82,7 +58,7 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae
text_encoder_two,
vae,
unet,
)
), sd_model_version
def download_weights(url, dest):
start = time.time()
+42 -3
View File
@@ -13,9 +13,12 @@ def get_unet_optimizer(
):
## unet_trainable_params can be unet.parameters() or a list of lora params
# 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(
@@ -35,6 +38,39 @@ def get_unet_optimizer(
print(f"Created {optimizer_name} optimizer for unet!")
return optimizer_unet
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
def is_belong_to_blocks(key, blocks):
try:
for g in blocks:
if g in key:
return True
return False
except Exception as e:
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
if use_blora:
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
try:
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
attns = [
attn_processor_name.rsplit(".", 1)[0]
for attn_processor_name, _ in unet.attn_processors.items()
if is_belong_to_blocks(attn_processor_name, blocks)
]
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
return target_modules
except Exception as e:
raise type(e)(
f"failed to get_target_modules, due to: {e}. "
f"Please check the modules specified in --lora_unet_blocks are correct"
)
def get_unet_lora_parameters(
lora_rank,
lora_alpha_multiplier: float,
@@ -43,12 +79,15 @@ def get_unet_lora_parameters(
unet,
pipe,
):
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
unet_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian",
target_modules=["to_k", "to_q", "to_v", "to_out.0", "conv2"],
#target_modules=["conv1", "conv2", "norm1", "norm2", "proj_in"], # TODO grid-search params for sd15
target_modules=target_modules,
use_dora=use_dora,
)
+48 -27
View File
@@ -1,7 +1,3 @@
# Have SwinIR upsample
# Have BLIP auto caption
# Have CLIPSeg auto mask concept
import gc
import fnmatch
import mimetypes
@@ -25,6 +21,7 @@ import numpy as np
import pandas as pd
import torch
from tqdm import tqdm
from transformers import (
BlipForConditionalGeneration,
Blip2ForConditionalGeneration,
@@ -38,13 +35,13 @@ from transformers import (
from trainer.utils.io import download_and_prep_training_data
from trainer.utils.utils import fix_prompt
from trainer.config import model_paths
import re
import openai
from openai import OpenAI
from dotenv import load_dotenv
load_dotenv()
try:
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
client = OpenAI(api_key=OPENAI_API_KEY)
@@ -54,8 +51,6 @@ except:
client = None
print("WARNING: Could not find OPENAI_API_KEY in .env, disabling gpt prompt generation.")
MODEL_PATH = "./cache"
# Put some boundaries to make the gpt pass work well: (very long text often confuses the model and also costs more money...)
MIN_GPT_PROMPTS = 3
MAX_GPT_PROMPTS = 50
@@ -139,7 +134,7 @@ def swin_ir_sr(
"""
model = Swin2SRForImageSuperResolution.from_pretrained(
model_id, cache_dir=MODEL_PATH
model_id, cache_dir = model_paths.get_path("SR")
).to(device)
processor = Swin2SRImageProcessor()
@@ -193,9 +188,9 @@ def clipseg_mask_generator(
model = None
if any(target_prompts):
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
model = CLIPSegForImageSegmentation.from_pretrained(
model_id, cache_dir=MODEL_PATH
model_id, cache_dir = model_paths.get_path("CLIP")
).to(device)
masks = []
@@ -408,14 +403,14 @@ def blip_caption_dataset(
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "blip2" in model_id:
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = Blip2ForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device)
else:
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
model = BlipForConditionalGeneration.from_pretrained(
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
).to(device)
for i, image in enumerate(tqdm(images)):
@@ -473,7 +468,7 @@ def gpt4_v_get_description(config, images):
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
payload = {
"model": "gpt-4-turbo",
"model": "gpt-4o",
"messages": [
{
"role": "user",
@@ -510,7 +505,7 @@ def gpt4_v_caption_dataset(
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
payload = {
"model": "gpt-4-turbo",
"model": "gpt-4o",
"messages": [
{
"role": "user",
@@ -643,6 +638,30 @@ def augment_image(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.
@@ -661,8 +680,6 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
return [new_width, new_height]
def load_and_save_masks_and_captions(
config,
concept_mode: str,
@@ -777,7 +794,10 @@ def load_and_save_masks_and_captions(
# Cleanup prompts using chatgpt:
captions = [fix_prompt(caption) for caption in captions]
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
trigger_text = ""
gpt_concept_description = None
if not config.disable_ti:
captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed)
aug_imgs, aug_caps = [],[]
# if we still have a very small amount of imgs, do some basic augmentation:
@@ -854,16 +874,11 @@ def load_and_save_masks_and_captions(
os.remove(os.path.join(output_dir, file))
os.makedirs(output_dir, exist_ok=True)
# Make sure we've correctly inserted the TOK into every caption:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
for caption in captions:
print(caption)
if config.remove_ti_token_from_prompts:
if config.disable_ti:
print('------------------ WARNING -------------------')
print("Removing 'TOK, ' from captions...")
print("This will completely break textual_inversion!!")
print("This will completely disable textual_inversion!!")
print('------------------ WARNING -------------------')
if gpt_concept_description:
replace_str = gpt_concept_description
@@ -871,6 +886,12 @@ def load_and_save_masks_and_captions(
replace_str = ""
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
captions = [caption.replace("TOK", replace_str) for caption in captions]
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
print("Final captions:")
for caption in captions:
print(caption)
# iterate through the images, masks, and captions and add a row to the dataframe for each
print("Saving final training dataset...")
+289
View File
@@ -0,0 +1,289 @@
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
+12 -3
View File
@@ -100,15 +100,17 @@ def print_system_info():
# Print disk space information
disk_usage = psutil.disk_usage('/')
free_disk = disk_usage.free // (1024 * 1024)
total_disk = disk_usage.total // (1024 * 1024)
used_disk = disk_usage.used // (1024 * 1024)
percent_disk_used = disk_usage.percent
print(f"Free disk space: {free_disk} MB with {percent_disk_used}% used")
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
# Print RAM information
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} MB with {percent_ram_used}% used")
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)}')
@@ -122,6 +124,13 @@ def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1
# Flatten and concatenate all parameters into a single tensor
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