Author SHA1 Message Date
mayukhdeb 6e4c86ea0c keep a copy 2024-07-22 11:07:14 -07:00
mayukhdeb 7e30d33dc2 find image filenames recursively in folder + train on big style dataset 2024-07-21 03:56:42 -07:00
mayukhdeb fbd050cb0c cleaner preprocess fn 2024-07-21 03:25:07 -07:00
mayukhdeb 70c8e88195 migrate to huggingface script 2024-07-21 02:05:30 -07:00
mayukhdeb 2f6f0cacd2 ignore stuff 2024-07-21 02:03:42 -07:00
mayukhdeb e8932f30ef start with a different embed string 2024-07-18 09:57:00 -07:00
mayukhdeb b7dd176536 add adamw8bit 2024-07-15 07:48:07 -07:00
mayukhdeb 5a72aca596 useful wandb name 2024-07-15 07:47:34 -07:00
mayukhdeb 85264bd004 wandb 2024-07-15 07:38:27 -07:00
mayukhdeb d46976f3dd sweep params set 2024-07-15 07:33:45 -07:00
mayukhdeb 8ed5996c1f train and inference on same device 2024-07-15 07:32:27 -07:00
mayukhdeb 8247b4b9dd fix OOM during inference (vae.decode) 2024-07-15 07:20:13 -07:00
mayukhdeb dc28fcb988 inference on same device 2024-07-15 03:28:31 -07:00
mayukhdeb 3c0a3be55f save different sh files for each gpu 2024-07-15 03:28:20 -07:00
mayukhdeb 274ec5c90b disable textual inversion if config.ti_lr is None 2024-07-11 12:46:52 -07:00
mayukhdeb fc54948eff more tweaks 2024-07-11 12:39:30 -07:00
mayukhdeb 32a7ae248d prompts for sweep 2024-07-11 12:34:42 -07:00
mayukhdeb 3cbabd2634 sweep dry runs 2024-07-11 12:26:07 -07:00
mayukhdeb 1cdacbc63a remove old todos 2024-07-11 11:16:22 -07:00
mayukhdeb 4668f755ec switch to adamw_8bit for sd3 transformer lora params 2024-07-11 00:15:09 -07:00
mayukhdeb dd093616b5 fix inference prompts bug 2024-07-09 09:09:56 -07:00
mayukhdeb 2c824d6807 more progress 2024-07-09 05:35:21 -07:00
mayukhdeb 022d51f53a inference fixed 2024-07-09 00:25:02 -07:00
mayukhdeb ecacf815ec re-impl textual inversion for first 2 text encoders 2024-07-08 12:45:55 -07:00
mayukhdeb 90d2269572 testing textual inversion training with frozen transformer 2024-07-04 07:40:41 -07:00
mayukhdeb 4955261cae new checkpoint + deterministic inference 2024-07-03 04:28:06 -07:00
mayukhdeb e84932af29 some small changes to text with main_sd3.py 2024-07-03 04:06:21 -07:00
mayukhdeb 115da83e2d cleaner output dir with checkpoints and generated samples in one folder 2024-07-03 03:43:30 -07:00
mayukhdeb 98b78cd6d6 cleanup + save training samples in output dir 2024-07-03 03:11:56 -07:00
mayukhdeb 7c5e0949ba keep changes 2024-07-02 05:19:34 -07:00
mayukhdeb ba2b7532ad more tweaks 2024-07-02 05:16:22 -07:00
mayukhdeb 2c0939a733 run inference less often 2024-07-02 04:25:26 -07:00
mayukhdeb 3b41479aa1 bfloat16 training + inference 2024-07-02 04:12:23 -07:00
mayukhdeb 6e589c76a8 impement some upstream changes and save a sample every 10 train steps 2024-07-02 01:50:14 -07:00
mayukhdeb 95f78d7b91 completely comement out TI for now 2024-07-01 01:56:49 -07:00
mayukhdeb 3c0ebd7d70 apply mask to loss + some hardcoding for banny debugging 2024-06-28 04:38:43 -07:00
mayukhdeb 6350b5344b more progress 2024-06-28 03:22:49 -07:00
mayukhdeb 2c4bd43044 temporarily remove ti grad norms 2024-06-28 03:12:16 -07:00
mayukhdeb b850a1615d watch grad norms 2024-06-28 03:04:40 -07:00
mayukhdeb 6d6dd98bfa dynamic ti lr 2024-06-28 02:03:34 -07:00
mayukhdeb 0444e35729 clip grad norms 2024-06-26 03:01:34 -07:00
mayukhdeb 0748e1d14e better prompt 2024-06-22 02:32:19 -07:00
mayukhdeb db7507c849 handle T5EncoderModel 2024-06-22 01:38:48 -07:00
mayukhdeb b379a28715 update todos 2024-06-22 01:38:12 -07:00
mayukhdeb 54ff8c4977 sd3 concept inference 2024-06-22 01:25:16 -07:00
mayukhdeb 64f28c8589 save TI embeds and lora adapters 2024-06-22 01:24:30 -07:00
mayukhdeb b12ce26fc9 update command 2024-06-20 03:52:30 -07:00
mayukhdeb b3da65dd39 ignore wandb stuff 2024-06-20 03:46:26 -07:00
mayukhdeb f64da5d4f1 update todo 2024-06-20 03:44:42 -07:00
mayukhdeb 5cc09092a4 tweak param 2024-06-20 03:43:44 -07:00
mayukhdeb 3d3d0ee4cb smash more todos 2024-06-20 03:42:32 -07:00
mayukhdeb 4b1efce0e3 compute loss and update weights 2024-06-20 03:28:49 -07:00
mayukhdeb d115c52274 do just forward passes 2024-06-20 03:20:41 -07:00
mayukhdeb 22127c9917 typo 2024-06-20 01:41:28 -07:00
mayukhdeb 4844845d5c more progress 2024-06-20 01:40:59 -07:00
mayukhdeb 3b15250bbc init train dataloader + update todos for training 2024-06-20 00:54:41 -07:00
mayukhdeb fd448e5437 accomodate T5EncoderModel 2024-06-20 00:54:21 -07:00
mayukhdeb a41cce1486 small cleanup 2024-06-20 00:27:51 -07:00
mayukhdeb 12b5960c36 progress bar for latent caching 2024-06-20 00:27:21 -07:00
mayukhdeb 2c2bfcda13 init PreprocessedDataset 2024-06-20 00:27:01 -07:00
mayukhdeb cacd6f2201 count trainable params from model 2024-06-19 23:58:26 -07:00
mayukhdeb 67024b4de6 full or lora finetuning of sd3 transformer 2024-06-19 23:58:16 -07:00
mayukhdeb 6d0d96ec79 more progress on todos 2024-06-19 23:37:31 -07:00
mayukhdeb fb2449a60d init textual inversion token embeds 2024-06-17 07:31:54 -07:00
mayukhdeb 0496ccdbef handle sd3 t5 text encoder 2024-06-17 07:31:33 -07:00
mayukhdeb 044aed9f03 sd3 train script wip 2024-06-17 06:42:16 -07:00
mayukhdeb bdda796c56 ignore notebook checkpoint 2024-06-17 05:21:54 -07:00
29 changed files with 2662 additions and 916 deletions
+9 -4
View File
@@ -1,16 +1,19 @@
data/
sd3_sweep_vis/
sd3_sweep_commands/
.ipynb_checkpoints/
cache
__pycache__
.ipynb_checkpoints/
models
lora_models*
eden_lora_training_runs/
datasets
*.tar
.env
.cog
.huggingface
train.py
rendered_images*
gridsearch*
@@ -21,4 +24,6 @@ conditioning_spaces/
training_args_x_*.json
xander_configs/
debug/*
wandb/
sd3_sweep_outputs/
sd3_face_sweep_configs/
+3 -4
View File
@@ -29,7 +29,7 @@ Install all dependencies using
then you can simply run:
`python main.py train_configs/training_args.json`
`python main.py -c training_args.json`
to start a training job.
Adjust the arguments inside `training_args.json` to setup a custom training job.
@@ -44,9 +44,8 @@ sudo curl -o /usr/local/bin/cog -L "https://github.com/replicate/cog/releases/la
sudo chmod +x /usr/local/bin/cog
```
2. Build the image with `cog build`
3. Run a training run with `sh cog_test_train.sh`
4. You can also go into the container with `cog run /bin/bash`
2. Build the image with `sudo cog build`
3. Run a training run with `sudo sh cog_test_train.sh`
## Automatic Checkpoint Evaluation
+7 -4
View File
@@ -3,14 +3,17 @@
build:
gpu: true
cuda: "12.1"
python_version: "3.11"
cuda: "11.8"
python_version: "3.9"
system_packages:
- "ffmpeg"
- "libgl1-mesa-glx"
- "libegl1-mesa-dev"
- "libsm6"
- "libxext6"
python_requirements: requirements.txt
run:
- wget https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task -O face_landmarker_v2_with_blendshapes.task
- wget http://thegiflibrary.tumblr.com/post/11565547760 -O face_landmarker_v2_with_blendshapes.task -q https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/1/face_landmarker.task
predict: "predict.py:Predictor"
image: "r8.im/edenartlab/sdxl-lora-trainer"
+1 -1
View File
@@ -1,5 +1,5 @@
# Set GPU ID to run these jobs on:
GPU_ID="device=3"
GPU_ID="device=2"
cog predict --gpus $GPU_ID \
-i name="xander_sdxl_cog" \
+179
View File
@@ -0,0 +1,179 @@
from trainer.utils.json_stuff import save_as_json
import itertools
import copy
import os
import random
random.seed(0)
GPU_IDS = [1,2,3]
wandb_log = True
def divide_list(lst, n):
"""
Divide a list into N equal parts.
Parameters:
lst (list): The list to be divided.
n (int): The number of parts to divide the list into.
Returns:
list of lists: A list containing N sublists, each of which is a part of the original list.
"""
if n <= 0:
raise ValueError("Number of parts must be greater than 0.")
if n > len(lst):
raise ValueError("Number of parts cannot be greater than the length of the list.")
# Calculate the size of each part
k, m = divmod(len(lst), n)
# Create the divided parts
return [lst[i * k + min(i, m):(i + 1) * k + min(i + 1, m)] for i in range(n)]
def generate_sh_file(commands, filename="script.sh"):
"""
Generates a .sh file with each command from the list written on a new line.
:param commands: List of commands to be written to the .sh file.
:param filename: Name of the .sh file to be created. Default is 'script.sh'.
"""
with open(filename, 'w') as file:
for command in commands:
file.write(command + '\n')
print(f"Saved: {filename}")
run_commands_dir = f"./sd3_sweep_commands"
os.system(
f"rm -rf {run_commands_dir} && mkdir -p {run_commands_dir}"
)
config_folder = "./sd3_face_sweep_configs"
os.system(f"rm -rf {config_folder}")
os.system(f"mkdir -p {config_folder}")
sweep_params = {
"unet_learning_rate": [
5e-5,
1e-4,
3e-4,
7e-4,
1e-3,
2e-3,
],
"train_batch_size": [
2,
4,
8,
16
],
"lora_rank": [
2,
4,
6,
8,
],
"ti_lr": [1e-3, None],
"unet_optimizer_type": [
"adamw",
"adamw_8bit",
"prodigy"
],
}
num_total_runs = 1
for key in sweep_params:
num_total_runs *= len(sweep_params[key])
print(f"Num total runs: {num_total_runs}")
default_config = {
"output_dir": "lora_models/sweep",
"sd_model_version": "sd3",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
"concept_mode": "face",
"seed": 0,
"resolution": 512,
"train_batch_size": 2,
"n_sample_imgs": 6,
"max_train_steps": 1000,
"token_warmup_steps": 200,
"checkpointing_steps": 1000, ## no need to save any checkpoints
"gradient_accumulation_steps": 2,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": False,
"text_encoder_lora_optimizer": None,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 16,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": False,
"caption_model": "blip",
"debug": True,
}
keys, values = zip(*sweep_params.items())
combinations = [dict(zip(keys, combination)) for combination in itertools.product(*values)]
all_config_paths = []
for index, c in enumerate(combinations):
config = copy.deepcopy(default_config)
filename = f"{index}"
# override default values with sweep params
for key in c:
"""
instead of editing the train batch size, we simply change the gradient
accumulation value. Which has the same effect.
We will also
"""
if key == "train_batch_size":
config["gradient_accumulation_steps"] = c[key] / config["train_batch_size"]
config["max_train_steps"] = config["max_train_steps"] * config["gradient_accumulation_steps"]
config["checkpointing_steps"] = config["checkpointing_steps"] * config["gradient_accumulation_steps"]
else:
config[key] = c[key]
# print(f"{index} - Setting {key} to {c[key]}")
filename += f"_{key}_{c[key]}"
config_path = os.path.join(
config_folder,
f"{filename}.json"
)
save_as_json(
dictionary_or_list=config,
filename = config_path
)
all_config_paths.append(config_path)
print(f"Saved: {config_path}")
print(f"Total: {index+1} configs")
all_commands = []
for c in all_config_paths:
command = f"python3 main_sd3.py {c}"
if wandb_log:
command = command + " --wandb-log"
all_commands.append(command)
random.shuffle(all_commands)
all_commands_split_by_gpu = divide_list(
lst = all_commands,
n = len(GPU_IDS)
)
for index, gpu_id in enumerate(GPU_IDS):
commands_on_single_gpu = [
f"CUDA_VISIBLE_DEVICES={gpu_id} {x}" for x in all_commands_split_by_gpu[index]
]
generate_sh_file(
commands = commands_on_single_gpu,
filename = os.path.join(
run_commands_dir,
f"run_on_gpu_{gpu_id}.sh"
)
)
+69 -178
View File
@@ -12,6 +12,7 @@ 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
@@ -25,7 +26,6 @@ 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,28 +34,10 @@ from trainer.optimizer import (
get_unet_optimizer
)
def train(config: TrainingConfig):
def train(
config: TrainingConfig,
):
seed_everything(config.seed)
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
), sd_model_version = load_models(config.pretrained_model, config.device, weight_dtype)
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,
@@ -76,6 +58,19 @@ def train(config: TrainingConfig):
if config.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
weight_dtype = dtype_map[config.weight_type]
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(config.pretrained_model, config.device, weight_dtype, keep_vae_float32=0)
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
text_encoders = [text_encoder_one, text_encoder_two],
@@ -118,26 +113,22 @@ def train(config: TrainingConfig):
embedding_handler.make_embeddings_trainable()
if not config.disable_ti:
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
else:
optimizer_ti = None
textual_inversion_params = None
optimizer_ti, textual_inversion_params = get_textual_inversion_optimizer(
text_encoders=text_encoders,
textual_inversion_lr=config.ti_lr,
textual_inversion_weight_decay=config.ti_weight_decay,
optimizer_name=config.ti_optimizer ## hardcoded
)
if not config.is_lora: # This code pathway has not been tested in a long while
print(f"Doing full fine-tuning on the U-Net")
unet.requires_grad_(True)
unet_lora_parameters = None
optimizer_text_encoder_lora = None
unet_trainable_params = unet.parameters()
else:
# Do lora-training instead.
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
# target_blocks=["block"] for original IP-Adapter
# target_blocks=["up_blocks.0.attentions.1"] for style blocks only
# target_blocks = ["up_blocks.0.attentions.1", "down_blocks.2.attentions.1"] # for style+layout blocks
@@ -203,7 +194,7 @@ def train(config: TrainingConfig):
print(f"--- Instantaneous batch size per device = {config.train_batch_size}")
print(f"--- Total batch_size (distributed + accumulation) = {total_batch_size}")
print(f"--- Gradient Accumulation steps = {config.gradient_accumulation_steps}")
print(f"--- Total optimization steps = {config.max_train_steps}\n", flush = True)
print(f"--- Total optimization steps = {config.max_train_steps}\n")
global_step = 0
last_save_step = 0
@@ -225,12 +216,10 @@ def train(config: TrainingConfig):
# 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
# let textual_inversion do the work first!
base_lr = 0.5e-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)
# let lora training kick in soonish
base_lr = 1.0e-4
#######################################################################################################
@@ -320,109 +309,7 @@ def train(config: TrainingConfig):
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())
@@ -433,7 +320,7 @@ def train(config: TrainingConfig):
loss += 0.0 * concept_description_loss
losses['concept_description_loss'].append(concept_description_loss.item())
if config.l1_penalty > 0.0 and unet_lora_parameters:
if config.l1_penalty > 0.0:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += config.l1_penalty * l1_norm
@@ -442,7 +329,6 @@ def train(config: TrainingConfig):
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()
@@ -463,6 +349,12 @@ def train(config: TrainingConfig):
grad_norms[f'text_encoder_{i}'].append(text_encoder_norm)
optimizer_collection.step()
# after every optimizer step, we do some manual intervention of the embeddings to regularize them:
if optimizer_collection.get_lr('textual_inversion') > 0.0:
#embedding_handler.fix_embedding_std(config.off_ratio_power)
pass
optimizer_collection.zero_grad()
#############################################################################################################
@@ -477,7 +369,7 @@ def train(config: TrainingConfig):
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:
if config.debug and (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > -1:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
@@ -498,29 +390,27 @@ def train(config: TrainingConfig):
)
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')
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')
embedding_handler.print_token_info()
plot_torch_hist(unet_lora_parameters if config.is_lora else unet.parameters(), global_step, config.output_dir, "lora_weights", min_val=-0.4, max_val=0.4, ymax_f = 0.08)
plot_loss(losses, save_path=f'{config.output_dir}/losses.png')
target_std_dict = {f"text_encoder_{idx}_target": embedding_handler.embeddings_settings[f"std_token_embedding_{idx}"].item() for idx in range(len(text_encoders)) if text_encoders[idx] is not None}
plot_token_stds(token_stds, save_path=f'{config.output_dir}/token_stds.png', target_value_dict=target_std_dict)
plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png')
plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png')
plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png')
validation_prompts = render_images(
pipe = pipe,
render_size = config.validation_img_size,
@@ -543,14 +433,14 @@ def train(config: TrainingConfig):
images_done += config.train_batch_size
global_step += 1
if global_step % (config.max_train_steps//50) == 0:
if global_step % (config.max_train_steps//20) == 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)
print_system_info()
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
yield np.min((progress, 1.0))
if global_step > config.max_train_steps:
print("Reached max steps, stopping training!", flush = True)
print("Reached max steps, stopping training!")
break
# final_save
@@ -581,7 +471,8 @@ def train(config: TrainingConfig):
pretrained_model_version=config.pretrained_model["version"]
)
if config.debug and 0:
print("Running final inference round...")
if config.debug:
# Reload the entire pipe from disk + LoRa:
pipe_to_use = None
checkpoint_folder = output_save_dir
@@ -620,6 +511,13 @@ def train(config: TrainingConfig):
img_grid_path = make_validation_img_grid(output_save_dir)
shutil.copy(img_grid_path, os.path.join(os.path.dirname(output_save_dir), f"validation_grid_{global_step:04d}.jpg"))
# Remove unneeded checkpoints if they exist in the output directory:
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_save_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
else:
print(f"Skipping final save, {output_save_dir} already exists")
@@ -633,8 +531,6 @@ def train(config: TrainingConfig):
config.job_time = time.time() - config.start_time
config.training_attributes["validation_prompts"] = validation_prompts
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
print("Training job complete, saving outputs...", flush = True)
print("------------------------------------------")
return config, output_save_dir
@@ -645,11 +541,6 @@ 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")
+2027
View File
File diff suppressed because it is too large Load Diff
+43 -58
View File
@@ -1,17 +1,22 @@
import os
import shutil
import tarfile
import json
import time
import random
import torch
import numpy as np
from PIL import Image
import pandas as pd
from dotenv import load_dotenv
from main import train
from trainer.config import TrainingConfig, model_paths
from trainer.utils.io import clean_filename
import folder_paths
import comfy.utils
from trainer.preprocess import preprocess
from trainer.models import pretrained_models
from trainer.config import TrainingConfig
from trainer.utils.io import clean_filename
from trainer.utils.utils import seed_everything
class Eden_LoRa_trainer:
@classmethod
@@ -19,9 +24,9 @@ class Eden_LoRa_trainer:
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"], ),
"lora_name": ("STRING", {"default": ""}),
"sd_model_version": (["sdxl", "sd15"], ),
"seed": ("INT", {"default": 0, "min": 0, "max": 100000}),
"resolution": ("INT", {"default": 512, "min": 256, "max": 768}),
"train_batch_size": ("INT", {"default": 4, "min": 1, "max": 8}),
"max_train_steps": ("INT", {"default": 400, "min": 50, "max": 1000}),
@@ -30,22 +35,17 @@ 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 = ("IMAGE", "STRING", "STRING", "STRING")
RETURN_NAMES = ("sample_images", "lora_path", "embedding_path", "final_msg")
RETURN_TYPES = ("STRING",)
FUNCTION = "train_lora"
def train_lora(self,
training_images_folder_path,
ckpt_name,
lora_name = "eden_lora",
mode = "style",
def train_lora(self, training_images_folder_path,
name = lora_name,
concept_mode = "style",
sd_model_version = "sdxl",
seed = 0,
resolution = 521,
train_batch_size = 4,
@@ -54,31 +54,21 @@ class Eden_LoRa_trainer:
unet_lr = 0.001,
lora_rank = 16,
use_dora = False,
n_tokens = 2,
debug_mode = False,
checkpointing_steps = 1000,
n_tokens = 2
):
print("Starting new training job...")
# Overwrite hardcoded paths to point to comfyUI folders:
model_paths.set_path("CLIP", os.path.join(folder_paths.models_dir, "clipseg"))
model_paths.set_path("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,
name="test",
lora_training_urls=training_images_folder_path,
concept_mode=mode,
ckpt_path=ckpt_path,
concept_mode=concept_mode,
sd_model_version=sd_model_version,
seed=seed,
resolution=resolution,
train_batch_size=train_batch_size,
max_train_steps=max_train_steps,
checkpointing_steps=checkpointing_steps,
checkpointing_steps=10000,
ti_lr=ti_lr,
unet_lr=unet_lr,
lora_rank=lora_rank,
@@ -86,45 +76,40 @@ class Eden_LoRa_trainer:
caption_model="blip",
n_tokens=n_tokens,
verbose=True,
debug=debug_mode,
debug=True,
)
pbar = comfy.utils.ProgressBar(100)
with torch.inference_mode(False):
train_generator = train(config=config)
while True:
try:
progress_f = next(train_generator)
pbar.update_absolute(progress_f * 100)
except StopIteration as e:
config, output_save_dir = e.value # Capture the return value
break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
out_path = f"{clean_filename(lora_name)}_eden_concept_lora_{int(time.time())}.tar"
directory = cogPath(output_save_dir)
with tarfile.open(out_path, "w") as tar:
print("Adding files to tar...")
for file_path in directory.rglob("*"):
print(file_path)
arcname = file_path.relative_to(directory)
tar.add(file_path, arcname=arcname)
# Add instructions README:
tar.add("instructions_README.md", arcname="README.md")
tar.add("comfyUI_workflow_lora_txt2img.json", arcname="comfyUI_workflow_lora_txt2img.json")
if sd_model_version == "sd15":
tar.add("comfyUI_workflow_lora_adiff.json", arcname="comfyUI_workflow_lora_adiff.json")
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
print(f"LORA training node finished in {config.job_time:.1f} seconds")
print("---------- Made with love by Eden.art 🌱 ----------")
# safetensors paths:
paths = [os.path.join(output_save_dir, f) for f in os.listdir(output_save_dir) if f.endswith(".safetensors")]
print(f"LORA training finished in {config.job_time:.1f} seconds")
print(f"Returning {out_path}")
# find the index of the path containing "_embeddings.safetensors":
for i, path in enumerate(paths):
if "_embeddings.safetensors" in path:
embedding_path = path
else:
lora_path = path
# Load the grid 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)
return (out_path,)
+16 -21
View File
@@ -1,22 +1,17 @@
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
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
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
numpy>=1.26.4
opencv-python>=4.1.0.25
mediapipe>=0.10.11
openai>=1.14.0
python-dotenv
prodigyopt
omegaconf
ujson
+20 -28
View File
@@ -35,12 +35,12 @@ def hamming_distance(dict1, dict2):
#######################################################################################
# Setup the base experiment config:
exp_name = "beeple"
exp_name = "grimes"
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
min_hamming_distance = 3 # min_n_params that have to be different from any previous experiment to be scheduled
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values
@@ -48,46 +48,43 @@ output_sh_path = f"gridsearch_configs/{exp_name}.sh"
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sdxl"],
"sd_model_version": ["sd15", "sdxl"],
"lora_training_urls": [
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large",
"/home/rednax/SSD2TB/Github_repos/Eden/images/beeple"
"/home/rednax/Documents/datasets/grimes"
],
"concept_mode": ['style'],
"sample_imgs_lora_scale": [0.8],
"disable_ti": ['false', 'true'],
"concept_mode": ['face'],
"seed": [0],
"resolution": [512],
"train_batch_size": [4],
"n_sample_imgs": [8],
"max_train_steps": [1200],
"checkpointing_steps": [200],
"n_sample_imgs": [6],
"max_train_steps": [400,800],
"checkpointing_steps": [100],
"gradient_accumulation_steps": [1],
"n_tokens": [2],
"ti_lr": [0.001],
"ti_weight_decay": [0.001],
"ti_lr": [0.001,0.0005],
"ti_weight_decay": [0.001,0.0],
"l1_penalty": [0.0],
"token_warmup_steps": [0],
"token_warmup_steps": [0,60],
"tok_cov_reg_w": [2000],
"cond_reg_w": [0.01e-5],
"tok_cond_reg_w": [0.01e-5],
"unet_lr": [0.0002, 0.00005],
"unet_prodigy_growth_factor": [1.05],
"unet_lr": [0.001],
"lora_alpha_multiplier": [1.0],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [16],
"use_dora": ['false'],
"unet_optimizer_type": ['AdamW8bit'],
"is_lora": ['false'],
"lora_rank": [16,32],
"use_dora": ['false', 'true'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["blip", "gpt4-v"],
"augment_imgs_up_to_n": [40],
"augment_imgs_up_to_n": [20,40],
"verbose": ['true'],
"debug": ['true']
}
@@ -149,12 +146,7 @@ def generate_sh_script(folder_path, output_sh_path):
# Write a command for each JSON file
for json_file in json_files:
file_path = os.path.join("scripts/", folder_path, json_file)
command = f"python main.py {file_path}\n"
if nohup:
command = f"nohup {command} > {file_path.replace('.json', '.log')} 2>&1 &\n"
command = f"python main.py {os.path.join(folder_path, json_file)}\n"
sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path)
+59
View File
@@ -0,0 +1,59 @@
"""
Pre-trained checkpoint:
https://huggingface.co/stabilityai/stable-diffusion-3-medium-diffusers
"""
import os
import torch
from diffusers import StableDiffusion3Pipeline
import torch
# Load the pretrained model
pipe = StableDiffusion3Pipeline.from_pretrained(
"stabilityai/stable-diffusion-3-medium-diffusers",
torch_dtype=torch.float16,
seed = 0
)
# Load the LoRA weights from file
lora_weights_path = "sd3-xander/checkpoint-1000/pytorch_lora_weights.safetensors"
# Move model to GPU
pipe = pipe.to("cuda")
prompts = [
"This is a picture of a man holding a glass of beer. He is wearing a casual plaid shirt and jeans. The man is holding a frosty glass of golden beer with a thick, foamy head in his right hand, lifting it slightly as if making a toast. The background features wooden tables and chairs, vintage beer signs, and warm ambient lighting",
"A close up shot of a man as a dragon rider with a red sword named Za'roc. His face is clearly visible in the high cinematic shot.",
"A man in 2075, looking for the last drop of water in mars. 4k HDR",
# "A king in Skyrim"
]
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}_baseline.jpg"
)
)
pipe.load_lora_weights(lora_weights_path, alpha = 8)
for idx, prompt in enumerate(prompts):
image = pipe(
prompt,
negative_prompt="",
num_inference_steps=28,
guidance_scale=7.0,
).images[0]
image.save(
os.path.join(
"./outputs",
f"{idx}.jpg"
)
)
print(f"Done!")
-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
-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
}
+3 -13
View File
@@ -135,7 +135,7 @@ def save_checkpoint(
embedding_handler.save_embeddings(
os.path.join(
output_dir,
f"{name}_{pretrained_model_version}_embeddings.safetensors"
f"{name}_embeddings.safetensors"
)
)
@@ -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,21 +184,11 @@ def save_checkpoint(
convert_pytorch_lora_safetensors_to_webui(
pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"),
output_filename=os.path.join(output_dir, f"{name}_{pretrained_model_version}_LoRa.safetensors")
output_filename=os.path.join(output_dir, f"{name}.safetensors")
)
else:
# Save the entire, finetuned unet weights:
unet.save_pretrained(save_directory = output_dir)
# Remove unneeded checkpoints if they exist in the output directory: TODO clean this up so they are never needed in the first place..
to_remove = ["pytorch_lora_weights.safetensors", "adapter_model.safetensors"]
for file in to_remove:
file_path = os.path.join(output_dir, file)
if os.path.exists(file_path):
os.remove(file_path)
return
def load_checkpoint(
pretrained_model_version: str,
pretrained_model_path: str,
+10 -48
View File
@@ -3,43 +3,15 @@ from datetime import datetime
from pydantic import BaseModel
import json, time, os
from typing import Literal
from trainer.models import pretrained_models
from trainer.utils.utils import pick_best_gpu_id
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", None] = None
ckpt_path: str = None # optional hardcoded checkpoint path
sd_model_version: Literal["sdxl", "sd15", "sd3"]
pretrained_model: dict = None
seed: Union[int, None] = None
resolution: int = 512
@@ -53,14 +25,14 @@ class TrainingConfig(BaseModel):
gradient_accumulation_steps: int = 1
is_lora: bool = True
unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw"
unet_optimizer_type: Literal["adamw", "prodigy", "adamw_8bit"] = "adamw"
unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer
unet_lr: float = 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
# if ti_lr is None, then we completely skip textual inversion
ti_lr: Union[float, None] = 1e-3
ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum
token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss
ti_weight_decay: float = 0.0
@@ -87,10 +59,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 = "eden_lora_training_runs"
output_dir: str = "lora_models/unnamed"
debug: bool = False
allow_tf32: bool = True
disable_ti: bool = False
remove_ti_token_from_prompts: bool = False
weight_type: Literal["fp16", "bf16", "fp32"] = "bf16"
n_tokens: int = 2
inserting_list_tokens: List[str] = ["<s0>","<s1>"]
@@ -102,7 +74,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 = None # Default lora scale for sampling the validation images
sample_imgs_lora_scale: float = 0.65 # Default lora scale for sampling the validation images
dataloader_num_workers: int = 0
training_attributes: dict = {}
aspect_ratio_bucketing: bool = False
@@ -121,11 +93,7 @@ class TrainingConfig(BaseModel):
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}
self.pretrained_model = pretrained_models[self.sd_model_version]
# add some metrics to the foldername:
lora_str = "dora" if self.use_dora else "lora"
@@ -134,7 +102,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"/{self.name}/" + f"{timestamp_short}-{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"--{timestamp_short}-{self.sd_model_version}_{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}"
os.makedirs(self.output_dir, exist_ok=True)
if self.seed is None:
@@ -148,12 +116,6 @@ 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.")
+2 -1
View File
@@ -6,6 +6,7 @@ import PIL
from PIL import Image
from torch.utils.data import Dataset
from typing import Tuple, Dict, List
from tqdm import tqdm
def prepare_image(
pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None,
@@ -68,7 +69,7 @@ class PreprocessedDataset(Dataset):
self.masks = []
self.do_cache = True
for idx in range(len(self.data)):
for idx in tqdm(range(len(self.data))):
if len(self.data) < 25:
print(self.captions[idx])
vae_latent, mask = self._process(idx)
+82 -23
View File
@@ -9,6 +9,7 @@ from typing import List, Optional, Dict
from safetensors.torch import save_file, safe_open
import matplotlib.pyplot as plt
from trainer.utils.utils import seed_everything, plot_torch_hist, plot_loss
from transformers import T5EncoderModel
class TokenEmbeddingsHandler:
def __init__(self, text_encoders, tokenizers):
@@ -31,7 +32,10 @@ class TokenEmbeddingsHandler:
continue
# Directly accessing and modifying the original weights tensor
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.requires_grad_(True)
else:
text_encoder.text_model.embeddings.token_embedding.weight.requires_grad_(True)
print(f"All embeddings in text_encoder_{idx} are now set to be trainable.")
def get_trainable_embeddings(self):
@@ -49,10 +53,22 @@ class TokenEmbeddingsHandler:
continue
# Ensure indices are a tensor. Use pre-existing dtype and device to match the model's.
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
if isinstance(text_encoder, T5EncoderModel):
indices_tensor = torch.tensor(
indices,
dtype=torch.long,
device=text_encoder.encoder.embed_tokens.weight.device
)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.encoder.embed_tokens.weight[indices_tensor]
else:
indices_tensor = torch.tensor(indices, dtype=torch.long, device=text_encoder.text_model.embeddings.token_embedding.weight.device)
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
# Directly access the embedding weights without detaching
token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight[indices_tensor]
embeddings[f'txt_encoder_{idx}'] = token_embeddings
# Get all corresponding tokens for these embeddings
@@ -192,9 +208,19 @@ class TokenEmbeddingsHandler:
self.non_train_ids = all_indices[inu]
# random initialization of new tokens
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
"""
handle both T5EncoderModel and other text encoders
T5EncoderModel is present in sd3
"""
if isinstance(text_encoder, T5EncoderModel):
std_token_embedding = (
text_encoder.encoder.embed_tokens.weight.data.std(dim=1).mean()
)
else:
std_token_embedding = (
text_encoder.text_model.embeddings.token_embedding.weight.data.std(dim=1).mean()
)
self.embeddings_settings[f"std_token_embedding_{idx}"] = std_token_embedding
if starting_toks is not None:
@@ -207,14 +233,28 @@ class TokenEmbeddingsHandler:
self.train_ids] = text_encoder.text_model.embeddings.token_embedding.weight.data[self.starting_ids].clone()
else:
std_multiplier = 1.0
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
init_embeddings = torch.randn(len(self.train_ids), text_encoder.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
else:
init_embeddings = torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype)
current_std = init_embeddings.std(dim=1).mean()
init_embeddings = init_embeddings * std_multiplier * std_token_embedding / current_std
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[self.train_ids] = init_embeddings.clone()
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids] = init_embeddings.clone()
if isinstance(text_encoder, T5EncoderModel):
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.encoder.embed_tokens.weight.data.clone()
else:
self.embeddings_settings[
f"original_embeddings_{idx}"
] = text_encoder.text_model.embeddings.token_embedding.weight.data.clone()
inu = torch.ones((len(tokenizer),), dtype=torch.bool)
inu[self.train_ids] = False
@@ -414,14 +454,27 @@ class TokenEmbeddingsHandler:
for idx, text_encoder in enumerate(self.text_encoders):
if text_encoder is None:
continue
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
if isinstance(text_encoder, T5EncoderModel):
assert text_encoder.encoder.embed_tokens.weight.data.shape[
0
] == len(self.tokenizers[idx]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
]
)
else:
assert text_encoder.text_model.embeddings.token_embedding.weight.data.shape[
0
] == len(self.tokenizers[0]), "Tokenizers should be the same."
new_token_embeddings = (
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
]
)
tensors[txt_encoder_keys[idx]] = new_token_embeddings
save_file(tensors, file_path)
@@ -474,9 +527,15 @@ class TokenEmbeddingsHandler:
self.train_ids = tokenizer.convert_tokens_to_ids(self.inserting_toks)
assert self.train_ids is not None, "New tokens could not be converted to IDs."
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
if isinstance(text_encoder, T5EncoderModel):
text_encoder.encoder.embed_tokens.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
else:
text_encoder.text_model.embeddings.token_embedding.weight.data[
self.train_ids
] = loaded_embeddings.to(device=self.device).to(dtype=self.dtype)
def load_embeddings(self, file_path: str, txt_encoder_keys = ["clip_l", "clip_g"]):
if not os.path.exists(file_path):
+14 -2
View File
@@ -4,6 +4,7 @@ import matplotlib.pyplot as plt
import torch
from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_foreach_support
from trainer.inference import get_conditioning_signals
from transformers import T5EncoderModel
def compute_snr(noise_scheduler, timesteps):
"""
@@ -104,7 +105,14 @@ class ConditioningRegularizer:
def __init__(self, config, embedding_handler):
self.config = config
self.embedding_handler = embedding_handler
self.target_norm = 34.5 if config.sd_model_version == 'sdxl' else 27.8
self.target_norms = {
"sdxl": 34.5,
"sd15": 27.8,
"sd3": 34.5
}
print(f'\033[91m[trainer.loss.ConditioningRegularizer] WARNING: Using a magic number: 34.5 for the target norm of sd3. We do not know if this is the ideal value. This might cause bugs or even break training completely.\033[0m')
self.target_norm = self.target_norms[config.sd_model_version]
self.reg_captions = ["a photo of TOK", "TOK", "a photo of TOK next to TOK", "TOK and TOK"]
self.token_replacement = config.token_dict.get("TOK", "TOK") # Fallback to "TOK" if not in dict
@@ -114,7 +122,11 @@ class ConditioningRegularizer:
if tokenizer is None:
idx += 1
continue
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
if isinstance(text_encoder, T5EncoderModel):
pretrained_token_embeddings = text_encoder.encoder.embed_tokens.weight.data
else:
pretrained_token_embeddings = text_encoder.text_model.embeddings.token_embedding.weight.data
self.distribution_regularizers[f'txt_encoder_{idx}'] = DistributionLoss(pretrained_token_embeddings, outdir = self.config.output_dir if config.debug else None)
idx += 1
+38 -12
View File
@@ -4,23 +4,49 @@ 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"
SD3_MODEL_CACHE = "models/stable-diffusion-3-medium"
#SD15_MODEL_CACHE = "./models/DreamShaper_6.31_BakedVae.safetensors"
#SD15_URL = "https://huggingface.co/Lykon/DreamShaper/resolve/main/DreamShaper_6.31_BakedVae.safetensors"
#SD15_MODEL_CACHE = "./models/photon_v1.safetensors"
#SD15_URL = "https://civitai.com/api/download/models/90072"
pretrained_models = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL, "version": "sdxl"},
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL, "version": "sd15"},
"sd3": {"path": SD3_MODEL_CACHE, "url": None, "version": "sd3"}
}
############################################################################################################
def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False):
if not isinstance(pretrained_model, dict) or 'path' not in pretrained_model or 'version' not in pretrained_model:
raise ValueError("pretrained_model must be a dict with 'path' and 'version' keys")
# check if the model is already downloaded:
if not os.path.exists(pretrained_model['path']):
download_weights(pretrained_model['url'], pretrained_model['path'])
print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...")
print(f"Loading model weights from {pretrained_model['path']} with dtype: {weight_dtype}...")
try:
pipe = StableDiffusionXLPipeline.from_single_file(
pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True)
sd_model_version = "sdxl"
except:
if pretrained_model['version'] == "sd15":
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!")
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)
@@ -36,14 +62,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 may not be ideal for training..?")
print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but might not be for training..")
unet.to(device, dtype=weight_dtype)
text_encoder_one.requires_grad_(False)
text_encoder_one.to(device, dtype=weight_dtype)
tokenizer_two = text_encoder_two = None
if sd_model_version == "sdxl":
if pretrained_model['version'] == "sdxl":
tokenizer_two = pipe.tokenizer_2
text_encoder_two = pipe.text_encoder_2
text_encoder_two.requires_grad_(False)
@@ -58,7 +84,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()
+8 -42
View File
@@ -3,6 +3,11 @@ import torch
import prodigyopt
from typing import Iterable
def count_trainable_params(model):
return sum([
x.numel() for x in model.parameters() if x.requires_grad
])
def get_unet_optimizer(
prodigy_d_coef: float,
prodigy_growth_factor: float,
@@ -13,12 +18,9 @@ def get_unet_optimizer(
):
## unet_trainable_params can be unet.parameters() or a list of lora params
# These learning rates will get overwritten in main.py:
if optimizer_name == "adamw":
optimizer_unet = torch.optim.AdamW(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
elif optimizer_name == "AdamW8bit":
import bitsandbytes as bnb
optimizer_unet = bnb.optim.AdamW8bit(unet_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay)
elif optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_unet = prodigyopt.Prodigy(
@@ -38,39 +40,6 @@ def get_unet_optimizer(
print(f"Created {optimizer_name} optimizer for unet!")
return optimizer_unet
# Taken (and slightly modified) from B-LoRA repo https://github.com/yardenfren1996/B-LoRA/blob/main/blora_utils.py
def is_belong_to_blocks(key, blocks):
try:
for g in blocks:
if g in key:
return True
return False
except Exception as e:
raise type(e)(f"failed to is_belong_to_block, due to: {e}")
def get_unet_lora_target_modules(unet, use_blora, target_blocks=None):
if use_blora:
content_b_lora_blocks = "unet.up_blocks.0.attentions.0"
style_b_lora_blocks = "unet.up_blocks.0.attentions.1"
target_blocks = [content_b_lora_blocks, style_b_lora_blocks]
try:
blocks = [(".").join(blk.split(".")[1:]) for blk in target_blocks]
attns = [
attn_processor_name.rsplit(".", 1)[0]
for attn_processor_name, _ in unet.attn_processors.items()
if is_belong_to_blocks(attn_processor_name, blocks)
]
target_modules = [f"{attn}.{mat}" for mat in ["to_k", "to_q", "to_v", "to_out.0", "conv2"] for attn in attns]
return target_modules
except Exception as e:
raise type(e)(
f"failed to get_target_modules, due to: {e}. "
f"Please check the modules specified in --lora_unet_blocks are correct"
)
def get_unet_lora_parameters(
lora_rank,
lora_alpha_multiplier: float,
@@ -79,15 +48,12 @@ def get_unet_lora_parameters(
unet,
pipe,
):
#target_modules = get_unet_lora_target_modules(unet, use_blora=True)
target_modules = ["to_k", "to_q", "to_v", "to_out.0", "conv2"]
unet_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank * lora_alpha_multiplier,
init_lora_weights="gaussian",
target_modules=target_modules,
target_modules=["to_k", "to_q", "to_v", "to_out.0", "conv2"],
#target_modules=["conv1", "conv2", "norm1", "norm2", "proj_in"], # TODO grid-search params for sd15
use_dora=use_dora,
)
+27 -48
View File
@@ -1,3 +1,7 @@
# Have SwinIR upsample
# Have BLIP auto caption
# Have CLIPSeg auto mask concept
import gc
import fnmatch
import mimetypes
@@ -21,7 +25,6 @@ import numpy as np
import pandas as pd
import torch
from tqdm import tqdm
from transformers import (
BlipForConditionalGeneration,
Blip2ForConditionalGeneration,
@@ -35,13 +38,13 @@ from transformers import (
from trainer.utils.io import download_and_prep_training_data
from trainer.utils.utils import fix_prompt
from trainer.config import model_paths
import re
import openai
from openai import OpenAI
from dotenv import load_dotenv
load_dotenv()
try:
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
client = OpenAI(api_key=OPENAI_API_KEY)
@@ -51,6 +54,8 @@ 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
@@ -134,7 +139,7 @@ def swin_ir_sr(
"""
model = Swin2SRForImageSuperResolution.from_pretrained(
model_id, cache_dir = model_paths.get_path("SR")
model_id, cache_dir=MODEL_PATH
).to(device)
processor = Swin2SRImageProcessor()
@@ -188,9 +193,9 @@ def clipseg_mask_generator(
model = None
if any(target_prompts):
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("CLIP"))
processor = CLIPSegProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = CLIPSegForImageSegmentation.from_pretrained(
model_id, cache_dir = model_paths.get_path("CLIP")
model_id, cache_dir=MODEL_PATH
).to(device)
masks = []
@@ -403,14 +408,14 @@ def blip_caption_dataset(
device=torch.device("cuda" if torch.cuda.is_available() else "cpu")
if "blip2" in model_id:
processor = Blip2Processor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
processor = Blip2Processor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = Blip2ForConditionalGeneration.from_pretrained(
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
).to(device)
else:
processor = BlipProcessor.from_pretrained(model_id, cache_dir = model_paths.get_path("BLIP"))
processor = BlipProcessor.from_pretrained(model_id, cache_dir=MODEL_PATH)
model = BlipForConditionalGeneration.from_pretrained(
model_id, cache_dir = model_paths.get_path("BLIP"), torch_dtype=torch.float16
model_id, cache_dir=MODEL_PATH, torch_dtype=torch.float16
).to(device)
for i, image in enumerate(tqdm(images)):
@@ -468,7 +473,7 @@ def gpt4_v_get_description(config, images):
base64_image = prep_img_for_gpt_api(img, max_size=(1024, 1024))
payload = {
"model": "gpt-4o",
"model": "gpt-4-turbo",
"messages": [
{
"role": "user",
@@ -505,7 +510,7 @@ def gpt4_v_caption_dataset(
base64_image = prep_img_for_gpt_api(img, max_size=(512, 512))
payload = {
"model": "gpt-4o",
"model": "gpt-4-turbo",
"messages": [
{
"role": "user",
@@ -638,30 +643,6 @@ 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.
@@ -680,6 +661,8 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
return [new_width, new_height]
def load_and_save_masks_and_captions(
config,
concept_mode: str,
@@ -794,10 +777,7 @@ def load_and_save_masks_and_captions(
# 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_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:
@@ -874,11 +854,16 @@ def load_and_save_masks_and_captions(
os.remove(os.path.join(output_dir, file))
os.makedirs(output_dir, exist_ok=True)
if config.disable_ti:
# Make sure we've correctly inserted the TOK into every caption:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
for caption in captions:
print(caption)
if config.remove_ti_token_from_prompts:
print('------------------ WARNING -------------------')
print("Removing 'TOK, ' from captions...")
print("This will completely disable textual_inversion!!")
print("This will completely break textual_inversion!!")
print('------------------ WARNING -------------------')
if gpt_concept_description:
replace_str = gpt_concept_description
@@ -886,12 +871,6 @@ def load_and_save_masks_and_captions(
replace_str = ""
captions = [caption.replace("TOK, ", replace_str + ", ") for caption in captions]
captions = [caption.replace("TOK", replace_str) for caption in captions]
else:
captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions]
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
@@ -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
+3 -12
View File
@@ -100,17 +100,15 @@ def print_system_info():
# Print disk space information
disk_usage = psutil.disk_usage('/')
total_disk = disk_usage.total // (1024 * 1024)
used_disk = disk_usage.used // (1024 * 1024)
free_disk = disk_usage.free // (1024 * 1024)
percent_disk_used = disk_usage.percent
print(f"Used disk space: {used_disk}/{total_disk} MB = {percent_disk_used}% used")
print(f"Free disk space: {free_disk} MB with {percent_disk_used}% used")
# Print RAM information
virtual_mem = psutil.virtual_memory()
total_ram = virtual_mem.total // (1024 * 1024)
current_ram = virtual_mem.used // (1024 * 1024)
percent_ram_used = virtual_mem.percent
print(f"Current used RAM: {current_ram}/{total_ram} MB = {percent_ram_used}% used")
print(f"Current used RAM: {current_ram} MB with {percent_ram_used}% used")
except Exception as e:
print(f'Error in gathering system info: {str(e)}')
@@ -124,13 +122,6 @@ def plot_torch_hist(parameters, step, checkpoint_dir, name, bins=100, min_val=-1
# Flatten and concatenate all parameters into a single tensor
all_params = torch.cat([p.data.view(-1) for p in parameters])
# count number of parameters:
n_params = len(all_params)
if n_params == 0 or n_params > 1e9:
return
norm = torch.norm(all_params)
# Convert to CPU for plotting
@@ -1,26 +1,29 @@
{
"name": "banny_sd15",
"output_dir": "lora_models/xander_sd15_final",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny.zip",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_big.zip",
"concept_mode": "face",
"sample_imgs_lora_scale": 0.8,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"train_batch_size": 7,
"n_sample_imgs": 6,
"max_train_steps": 800,
"max_train_steps": 5000,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"checkpointing_steps": 50,
"gradient_accumulation_steps": 1,
"sample_imgs_lora_scale": 0.8,
"n_tokens": 2,
"ti_lr": 0.001,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_lr": 0.5e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"text_encoder_lora_rank": 16,
"unet_lr": 0.001,
"lora_alpha_multiplier": 1.0,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
@@ -1,27 +1,25 @@
{
"name": "xander_test",
"output_dir": "lora_models/object",
"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",
"lora_training_urls": "/home/rednax/Documents/datasets/DOV/lizzo/full body",
"concept_mode": "object",
"seed": 1,
"resolution": 512,
"resolution": 640,
"train_batch_size": 4,
"n_sample_imgs": 4,
"max_train_steps": 200,
"max_train_steps": 420,
"token_warmup_steps": 0,
"checkpointing_steps": 100,
"checkpointing_steps": 60,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"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,
"lora_rank": 12,
"use_dora": false,
"caption_model": "blip",
"caption_model": "gpt4-v",
"debug": true
}
@@ -1,15 +1,17 @@
{
"name": "clipx_sd15",
"output_dir": "lora_models/does_best",
"sd_model_version": "sd15",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip",
"concept_mode": "style",
"seed": 0,
"resolution": 512,
"seed": 1,
"resolution": 640,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"max_train_steps": 600,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
@@ -1,28 +1,29 @@
{
"name": "clipx_sdxl",
"output_dir": "lora_models/Journey",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"lora_training_urls": "/home/rednax/Documents/datasets/journey",
"concept_mode": "style",
"sample_imgs_lora_scale": 0.7,
"seed": 1,
"seed": 0,
"resolution": 512,
"train_batch_size": 4,
"n_sample_imgs": 6,
"max_train_steps": 400,
"max_train_steps": 1000,
"token_warmup_steps": 0,
"checkpointing_steps": 200,
"checkpointing_steps": 100,
"gradient_accumulation_steps": 1,
"n_tokens": 2,
"ti_lr": 0.001,
"ti_weight_decay": 0.0005,
"remove_ti_token_from_prompts": false,
"text_encoder_lora_optimizer": null,
"text_encoder_lora_lr": 1.0e-4,
"text_encoder_lora_weight_decay": 1e-5,
"text_encoder_lora_rank": 12,
"unet_lr": 0.001,
"prodigy_d_coef": 1.0,
"unet_prodigy_growth_factor": 1.05,
"lora_rank": 16,
"use_dora": false,
"caption_model": "blip",
"use_dora": true,
"caption_model": "gpt4-v",
"debug": true
}