37 Commits
Author SHA1 Message Date
aiXander 8b340a6edc unpushed changes 2024-03-29 10:38:52 -07:00
aiXander 322b3da61e merge 2024-03-15 09:12:01 -07:00
aiXander 7da3e64b1e updates 2024-03-15 09:10:58 -07:00
mayukhdeb 5088ad3754 train_pti: switch to 500 steps + fix precision arg 2024-03-15 02:04:30 -07:00
mayukhdeb d5e56dc6e1 use keyword args for sanity 2024-03-15 02:04:02 -07:00
mayukhdeb 4563279eca config: forbid extras 2024-03-15 02:03:08 -07:00
mayukhdeb bf892826f7 oopsie 2024-03-15 00:48:47 -07:00
mayukhdeb 661f41c7c6 hardcode l1_penalty to be 0.05 if concept_mode == "style" 2024-03-15 00:48:14 -07:00
mayukhdeb fad22f61cc trainer_pti: keep all args in preprocess and init 2024-03-15 00:41:28 -07:00
mayukhdeb bbdd8d4762 trainer: fix prodigy_d_coef 2024-03-15 00:20:11 -07:00
mayukhdeb 42adc67c3a fix prodigy_d_coef default value 2024-03-15 00:19:43 -07:00
mayukhdeb f020bcf218 trainer: fix undefined variable 2024-03-14 22:04:34 -07:00
aiXander 838856fcef Merge branch 'main' of https://github.com/edenartlab/trainer into main 2024-03-14 21:32:25 -07:00
mayukhdeb 87fc89141b switch to fp32 2024-03-14 21:22:04 -07:00
aiXander 12bd221dbe trying to fix trainer 2024-03-14 21:05:46 -07:00
aiXander 8d2517741c more updates 2024-03-14 20:39:19 -07:00
aiXander 580d7867a8 make training work again 2024-03-14 20:12:30 -07:00
aiXander b35f4529dd fix more bugs 2024-03-14 19:45:35 -07:00
aiXander 28a9772cf6 large amounts of bugfixes 2024-03-14 19:24:24 -07:00
aiXander 0722ab6ce7 sync training args, add some bugfixes 2024-03-14 18:18:52 -07:00
aiXander 416a4622c2 add download weights 2024-03-14 17:45:56 -07:00
Xander Steenbrugge 40edbd0e85 add special_params.json and training_args.json 2024-03-14 17:05:26 -07:00
mayukhdeb 2cd456020b add concept_mode arg 2024-03-14 12:51:06 -07:00
mayukhdeb d4b550792d render_images regardless of debug or not 2024-03-14 12:50:42 -07:00
mayukhdeb 01d112f140 cleanup 2024-03-14 07:04:51 -07:00
mayukhdeb f87ab80048 save train args even for intermediate checkpoints 2024-03-14 06:59:29 -07:00
mayukhdeb 5ebc20ace0 return only output save dir 2024-03-14 06:57:36 -07:00
mayukhdeb 6b2f19ceb9 dont store validation prompts 2024-03-14 06:50:44 -07:00
mayukhdeb 53bda05d3d dont store validation_promts in train config 2024-03-14 06:50:16 -07:00
mayukhdeb b4a76f12ba stop ignoring trainer folder + remove train.py 2024-03-14 06:42:58 -07:00
mayukhdeb 9b2bc5da0f add trainer class 2024-03-14 06:42:58 -07:00
mayukhdeb 3387701f9a add config 2024-03-14 06:42:58 -07:00
mayukhdeb c9c9c89442 add utils 2024-03-14 06:42:58 -07:00
mayukhdeb 9de4a47c31 add init files 2024-03-14 06:42:58 -07:00
mayukhdeb c8bf5ce3b5 add prepare_prompt_for_lora 2024-03-14 06:42:58 -07:00
mayukhdeb a4ba97aebc move stuff into trainer module 2024-03-14 06:42:58 -07:00
mayukhdeb a03e1f6cb0 move files 2024-03-14 06:42:58 -07:00
17 changed files with 1122 additions and 880 deletions
-2
View File
@@ -9,7 +9,5 @@ __pycache__
xander*.sh xander*.sh
.huggingface .huggingface
tests/ tests/
trainer/
train.py
debug/* debug/*
!debug/*.py !debug/*.py
+2
View File
@@ -30,6 +30,7 @@ Algo:
- the random initialization of the token embeddings has a relatively large impact on the final outcome, there are prob ways to reduce - the random initialization of the token embeddings has a relatively large impact on the final outcome, there are prob ways to reduce
this random variance, eg CLIP_similarity pretraining. this random variance, eg CLIP_similarity pretraining.
- Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM - Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM
- it looks the like gpu-utilization is only like 65-70% during training: whats the bottleneck? Can we speed this up?
Bugfixing: Bugfixing:
see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037 see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037
@@ -50,6 +51,7 @@ Bigger improvements:
Tuning Experiments once code is fully ready: Tuning Experiments once code is fully ready:
- test if VAE weight_type actually matters for training
- try-out conditioning noise injection during training to increase robustness - try-out conditioning noise injection during training to increase robustness
- re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam) - re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam)
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible), - right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
+27 -4
View File
@@ -13,13 +13,36 @@ SDXL_MODEL_CACHE = "./models/juggernaut_v6.safetensors"
SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors" SDXL_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernautXL_v6.safetensors"
SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors" SD15_MODEL_CACHE = "./models/juggernaut_reborn.safetensors"
# TODO point this url to the correct full folder structure containing the CLIP text-encoder (this wont actually work rn)
SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors" SD15_URL = "https://edenartlab-lfs.s3.amazonaws.com/models/checkpoints/juggernaut_reborn.safetensors"
SDXL_TURBO_MODEL_CACHE = "./models/SDXL_turbo.safetensors"
SDXL_TURBO_URL = "https://huggingface.co/stabilityai/sdxl-turbo/resolve/main/sd_xl_turbo_1.0_fp16.safetensors?download=true"
SDXL_LIGHTNING_MODEL_CACHE = "./models/SDXL_lightning.safetensors"
SDXL_LIGHTNING_URL = "https://huggingface.co/ByteDance/SDXL-Lightning/resolve/main/sdxl_lightning_8step.safetensors?download=true"
# Define model paths and URLs in a dictionary # Define model paths and URLs in a dictionary
MODEL_INFO = { MODEL_DICT = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL}, "sdxl": {
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL} "path": SDXL_MODEL_CACHE,
"url": SDXL_URL,
"version": "sdxl"
},
"sd15": {
"path": SD15_MODEL_CACHE,
"url": SD15_URL,
"version": "sd15"
},
"sdxl_turbo": {
"path": SDXL_TURBO_MODEL_CACHE,
"url": SDXL_TURBO_URL,
"version": "sdxl_turbo"
},
"sdxl_lightning": {
"path": SDXL_LIGHTNING_MODEL_CACHE,
"url": SDXL_LIGHTNING_URL,
"version": "sdxl_lightning"
}
} }
def download_weights(url, dest): def download_weights(url, dest):
-89
View File
@@ -1,89 +0,0 @@
import os, json
import torch
from safetensors.torch import load_file
from typing import Dict
from peft import PeftModel
from dataset_and_utils import TokenEmbeddingsHandler
from safetensors.torch import save_file
'''
from diffusers.utils import (
convert_all_state_dict_to_peft,
convert_state_dict_to_diffusers,
convert_unet_state_dict_to_peft
)
'''
def patch_pipe_with_lora(pipe, lora_path):
"""
update the pipe with the lora model and the token embeddings
"""
pipe.unet = PeftModel.from_pretrained(pipe.unet, lora_path)
pipe.unet.merge_adapter()
# Load the textual_inversion token embeddings into the pipeline:
try: #SDXL
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
except: #SD15
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
embeddings_path = [f for f in os.listdir(lora_path) if f.endswith("embeddings.safetensors")][0]
handler.load_embeddings(os.path.join(lora_path, embeddings_path))
return pipe
def unet_attn_processors_state_dict(unet) -> Dict[str, torch.tensor]:
"""
Returns:
a state dict containing just the attention processor parameters.
"""
attn_processors = unet.attn_processors
attn_processors_state_dict = {}
for attn_processor_key, attn_processor in attn_processors.items():
for parameter_key, parameter in attn_processor.state_dict().items():
attn_processors_state_dict[
f"{attn_processor_key}.{parameter_key}"
] = parameter
return attn_processors_state_dict
def save_lora(output_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names):
"""
Save the LORA model to output_dir, optionally with some example images
"""
print(f"Saving checkpoint at step.. {global_step}")
os.makedirs(output_dir, exist_ok=True)
args_dict["n_training_steps"] = global_step
args_dict["total_n_imgs_seen"] = global_step * args_dict["train_batch_size"]
if not is_lora:
lora_tensors = {
name: param
for name, param in unet.named_parameters()
if name in unet_param_to_optimize_names
}
save_file(lora_tensors, f"{output_dir}/unet.safetensors",)
elif len(unet_lora_parameters) > 0:
unet.save_pretrained(save_directory = output_dir)
try:
concept_name = args_dict["name"].lower()
except:
concept_name = "eden_concept_lora"
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
concept_name = concept_name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
embedding_handler.save_embeddings(f"{output_dir}/{concept_name}_embeddings.safetensors",)
with open(f"{output_dir}/special_params.json", "w") as f:
json.dump(token_dict, f)
with open(f"{output_dir}/training_args.json", "w") as f:
json.dump(args_dict, f, indent=4)
+2
View File
@@ -0,0 +1,2 @@
from .trainer import Trainer
from .config import TrainerConfig
+61
View File
@@ -0,0 +1,61 @@
from typing import Optional, List, Dict, Any
from pydantic import BaseModel, Field
import random
import json
from typing import Literal
import torch
precision_map = {
"fp16": torch.float16,
"bf16": torch.bfloat16,
"fp32": torch.float32
}
class TrainerConfig(BaseModel, extra = "forbid"):
pretrained_model: Dict[str, str] # should be a dict with keys "path" and "version"
name: str='unnamed',
trigger_text: str='a photo of TOK, ',
instance_data_dir: str = "./dataset/zeke/captions.csv"
concept_mode: Literal["face", "concept", "object", "style"]
output_dir: str = "lora_output"
seed: Optional[int] = Field(default_factory=lambda: random.randint(0, 2**32 - 1))
resolution: int = 960
crops_coords_top_left_h: int = 0
crops_coords_top_left_w: int = 0
train_batch_size: int = 1
train_dataset_cache: bool = True
num_train_epochs: int = 10000
max_train_steps: Optional[int] = None
checkpointing_steps: int = 500000
gradient_accumulation_steps: int = 1
unet_learning_rate: float = 1.0
textual_inversion_lr: float = 1e-3
textual_inversion_weight_decay: float = 3e-4
prodigy_d_coef: float = 0.5,
l1_penalty: float = 0.0
lora_weight_decay: float = 0.005
scale_lr_based_on_grad_acc: bool = False
lr_scheduler_name: str = "constant"
lr_warmup_steps: int = 50
lr_num_cycles: int = 1
lr_power: float = 1.0
snr_gamma: float = 5.0
dataloader_num_workers: int = 0
allow_tf32: bool = True
precision: Literal["bf16", "fp16", "fp32"] = "bf16"
optimizer_name: Literal["prodigy", "adamw"] = "prodigy"
device: str = "cuda"
token_dict: Dict[str, str] = {"TOK": "<s0><s1>"}
inserting_list_tokens: List[str] = ["<s0><s1>"]
verbose: bool = True
is_lora: bool = True
lora_rank: int = 12
lora_alpha: int = 12
args_dict: Dict[str, Any] = {}
debug: bool = False
hard_pivot: bool = True
off_ratio_power: float = 0.1
def save_as_json(self, file_path: str) -> None:
with open(file_path, 'w') as f:
json.dump(self.dict(), f, indent=4)
@@ -35,7 +35,7 @@ def plot_torch_hist(parameters, epoch, checkpoint_dir, name, bins=100, min_val=-
plt.xlim(min_val, max_val) plt.xlim(min_val, max_val)
plt.xlabel('Weight Value') plt.xlabel('Weight Value')
plt.ylabel('Count') plt.ylabel('Count')
plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f})') plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f}, min = {np.min(all_params_cpu):.2f}, max = {np.max(all_params_cpu):.2f})')
plt.savefig(f"{checkpoint_dir}/{name}_histogram_{epoch:04d}.png") plt.savefig(f"{checkpoint_dir}/{name}_histogram_{epoch:04d}.png")
plt.close() plt.close()
@@ -486,7 +486,6 @@ class TokenEmbeddingsHandler:
print("Initializing new tokens...") print("Initializing new tokens...")
print(inserting_toks) print(inserting_toks)
torch.manual_seed(seed)
idx = 0 idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders): for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
@@ -541,6 +540,7 @@ class TokenEmbeddingsHandler:
else: else:
if 1: # random initialization: if 1: # random initialization:
torch.manual_seed(seed)
init_embeddings = (torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype) * std_token_embedding * 1.0) init_embeddings = (torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype) * std_token_embedding * 1.0)
else: else:
# Test code to initialize the new tokens with some specific tokens # Test code to initialize the new tokens with some specific tokens
+590
View File
@@ -0,0 +1,590 @@
import os
import math
import random
import numpy as np
import torch
import fnmatch
from peft import LoraConfig, get_peft_model
from diffusers.optimization import get_scheduler
from tqdm import tqdm
import shutil
import time
import gc
import prodigyopt
from .config import (
TrainerConfig,
precision_map
)
from .dataset_and_utils import (
load_models,
TokenEmbeddingsHandler,
PreprocessedDataset,
plot_torch_hist,
plot_loss,
plot_lrs
)
from .utils.model_info import print_trainable_parameters
from .utils.snr import compute_snr
from .utils.learning_rate import get_avg_lr
from .utils.lora import save_lora
from .utils.rendering import render_images
from io_utils import download_weights
from preprocess import preprocess
class Trainer:
def __init__(self, args):
self.args = args
random.seed(args.seed)
torch.manual_seed(args.seed)
np.random.seed(args.seed)
torch.cuda.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
#torch.backends.cudnn.deterministic = True
print("Trainer initialized!")
def train(self):
if self.args.concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
self.args.l1_penalty = 0.05
args = self.args
if args.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
weight_dtype = precision_map[args.precision]
print(f"Loading models with weight_dtype: {weight_dtype}")
if args.scale_lr_based_on_grad_acc:
unet_learning_rate = (
args.unet_learning_rate * args.gradient_accumulation_steps * args.train_batch_size
)
# Download the weights if they don't exist locally
if not os.path.exists(args.pretrained_model['path']):
download_weights(args.pretrained_model['url'], args.pretrained_model['path'])
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(
pretrained_model = args.pretrained_model,
device=args.device,
weight_dtype=weight_dtype
)
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
[text_encoder_one, text_encoder_two], [tokenizer_one, tokenizer_two]
)
starting_toks = None
embedding_handler.initialize_new_tokens(
inserting_toks=args.inserting_list_tokens,
starting_toks=starting_toks,
seed=args.seed
)
text_encoders = [text_encoder_one, text_encoder_two]
unet_param_to_optimize = []
text_encoder_parameters = []
for text_encoder in text_encoders:
if text_encoder is not None:
for name, param in text_encoder.named_parameters():
if "token_embedding" in name:
param.requires_grad = True
text_encoder_parameters.append(param)
else:
param.requires_grad = False
unet_param_to_optimize_names = []
unet_lora_parameters = []
if not args.is_lora:
WHITELIST_PATTERNS = [
# "*.attn*.weight",
# "*ff*.weight",
"*"
]
BLACKLIST_PATTERNS = ["*.norm*.weight", "*time*"]
for name, param in unet.named_parameters():
if any(
fnmatch.fnmatch(name, pattern) for pattern in WHITELIST_PATTERNS
) and not any(
fnmatch.fnmatch(name, pattern) for pattern in BLACKLIST_PATTERNS
):
param.requires_grad_(True)
unet_param_to_optimize_names.append(name)
print(f"Training: {name}")
else:
param.requires_grad_(False)
# Optimizer creation
params_to_optimize = [
{
"params": text_encoder_parameters,
"lr": args.textual_inversion_lr,
"weight_decay": args.textual_inversion_weight_decay,
},
]
params_to_optimize_prodigy = [
{
"params": unet_param_to_optimize,
"lr": unet_learning_rate,
"weight_decay": args.lora_weight_decay,
},
]
else:
# Do lora-training instead.
unet.requires_grad_(False)
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
use_dora = True
unet_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights="gaussian",
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
use_dora=use_dora,
)
if use_dora:
print(f"Disabling L1 penalty for DORA training")
args.l1_penalty = 0.0
unet = get_peft_model(unet, unet_lora_config)
print_trainable_parameters(unet, name = 'unet')
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
# Loop over the unet_lora_parameters and print their names and shapes:
for name, param in unet.named_parameters():
if param.requires_grad:
print(name, param.shape)
params_to_optimize = [{
"params": text_encoder_parameters,
"lr": args.textual_inversion_lr,
"weight_decay": args.textual_inversion_weight_decay,
}]
params_to_optimize_prodigy = [{
"params": unet_lora_parameters,
"lr": 1.0,
"weight_decay": args.lora_weight_decay,
}]
if args.optimizer_name == "adamw":
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=0.0, # this wd doesn't matter, I think
)
optimizer_prod = None
elif args.optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_prod = prodigyopt.Prodigy(
params_to_optimize_prodigy,
d_coef = args.prodigy_d_coef,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=args.lora_weight_decay,
betas=(0.9, 0.99),
growth_rate=1.025, # this slows down the lr_rampup
#growth_rate=1.05, # this slows down the lr_rampup
)
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=args.textual_inversion_weight_decay,
)
train_dataset = PreprocessedDataset(
args.instance_data_dir,
tokenizer_one,
tokenizer_two,
vae,
do_cache=args.train_dataset_cache,
substitute_caption_map=args.token_dict,
)
print(f"# PTI : Loaded dataset, do_cache: {args.train_dataset_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=args.train_batch_size,
shuffle=True,
num_workers=args.dataloader_num_workers,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps
)
if args.max_train_steps is None:
max_train_steps = num_train_epochs * num_update_steps_per_epoch
else:
max_train_steps = args.max_train_steps
lr_scheduler = get_scheduler(
args.lr_scheduler_name,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * args.gradient_accumulation_steps,
num_training_steps=max_train_steps * args.gradient_accumulation_steps,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps
)
num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch)
total_batch_size = args.train_batch_size * args.gradient_accumulation_steps
if args.verbose:
print(f"# PTI : Running training ")
print(f"# PTI : Num examples = {len(train_dataset)}")
print(f"# PTI : Num batches each epoch = {len(train_dataloader)}")
print(f"# PTI : Num Epochs = {num_train_epochs}")
print(f"# PTI : Instantaneous batch size per device = {args.train_batch_size}")
print(f"# PTI : Total train batch size (distributed & accumulation) = {total_batch_size}")
print(f"# PTI : Gradient Accumulation steps = {args.gradient_accumulation_steps}")
print(f"# PTI : Total optimization steps = {max_train_steps}")
global_step = 0
first_epoch = 0
last_save_step = 0
progress_bar = tqdm(range(global_step, max_train_steps), position=0, leave=True)
checkpoint_dir = os.path.join(args.output_dir, "checkpoints")
if os.path.exists(checkpoint_dir):
shutil.rmtree(checkpoint_dir)
os.makedirs(f"{checkpoint_dir}")
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
#embedding_handler.pre_optimize_token_embeddings(train_dataset)
ti_lrs, lora_lrs = [], []
losses = []
start_time, images_done = time.time(), 0
for epoch in range(first_epoch, num_train_epochs):
unet.train()
progress_bar.set_description(f"# PTI :step: {global_step}, epoch: {epoch}")
for step, batch in enumerate(train_dataloader):
progress_bar.update(1)
if args.hard_pivot:
if epoch >= num_train_epochs // 2:
if optimizer is not None:
print("----------------------")
print("# PTI : Pivot halfway")
print("----------------------")
# remove text encoder parameters from the optimizer
optimizer.param_groups = None
# remove the optimizer state corresponding to text_encoder_parameters
for param in text_encoder_parameters:
if param in optimizer.state:
del optimizer.state[param]
optimizer = None
else: # Update learning rates gradually:
finegrained_epoch = epoch + step / len(train_dataloader)
completion_f = finegrained_epoch / num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training
optimizer.param_groups[0]['lr'] = args.textual_inversion_lr * (1 - completion_f) ** 2.0
try: #sdxl
(tok1, tok2), vae_latent, mask = batch
except: #sd15
tok1, vae_latent, mask = batch
tok2 = None
vae_latent = vae_latent.to(weight_dtype)
# tokens to text embeds
prompt_embeds_list = []
for tok, text_encoder in zip((tok1, tok2), text_encoders):
if tok is None:
continue
prompt_embeds_out = text_encoder(
tok.to(text_encoder.device),
output_hidden_states=True,
)
pooled_prompt_embeds = prompt_embeds_out[0]
prompt_embeds = prompt_embeds_out.hidden_states[-2]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
prompt_embeds_list.append(prompt_embeds)
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
# Create Spatial-dimensional conditions.
original_size = (args.resolution, args.resolution)
target_size = (args.resolution, args.resolution)
crops_coords_top_left = (
args.crops_coords_top_left_h,
args.crops_coords_top_left_w
)
add_time_ids = list(original_size + crops_coords_top_left + target_size)
add_time_ids = torch.tensor([add_time_ids])
add_time_ids = add_time_ids.to(
args.device,
dtype=prompt_embeds.dtype
).repeat(
bs_embed, 1
)
# Sample noise that we'll add to the latents:
noise = torch.randn_like(vae_latent)
noise_offset = 0.05 # TODO, turn this into an input arg and do a grid search
if noise_offset > 0.0:
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
noise += noise_offset * torch.randn(
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
bsz = vae_latent.shape[0]
timesteps = torch.randint(
0,
noise_scheduler.config.num_train_timesteps,
(bsz,),
device=vae_latent.device,
).long()
noisy_model_input = noise_scheduler.add_noise(vae_latent, noise, timesteps)
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,1:-2,:]) * noise_sigma
# Predict the noise residual
model_pred = unet(
noisy_model_input,
timesteps,
prompt_embeds,
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
).sample
# Get the unet prediction target depending on the prediction type:
if noise_scheduler.config.prediction_type == "epsilon":
target = noise
else:
raise NotImplementedError(f"Not implemented for noise_scheduler.config.prediction_type: {noise_scheduler.config.prediction_type}")
# Compute the loss:
if args.snr_gamma is None:
loss = (model_pred - target).pow(2) * mask
# modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
# Average the normalized errors across the batch
loss = loss.mean()
else:
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
# This is discussed in Section 4.2 of the same paper.
snr = compute_snr(noise_scheduler, timesteps)
base_weight = (
torch.stack([snr, args.snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
)
if noise_scheduler.config.prediction_type == "v_prediction":
# Velocity objective needs to be floored to an SNR weight of one.
mse_loss_weights = base_weight + 1
else:
# Epsilon and sample both use the same loss weights.
mse_loss_weights = base_weight
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
loss = (model_pred - target).pow(2) * mask
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
if 1: # modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
loss = loss.mean()
if args.l1_penalty > 0.0:
# Compute normalized L1 norm (mean of abs sum) of all lora parameters:
l1_norm = sum(p.abs().sum() for p in unet_lora_parameters) / sum(p.numel() for p in unet_lora_parameters)
loss += args.l1_penalty * l1_norm
# Print the relative L1 norm:
if global_step % 50 == 0:
print(f" ---- L1 norm: {l1_norm.item():.4f}")
print(f" ---- L1 loss: {args.l1_penalty * l1_norm.item():.4f}")
print(f" ---- Total loss: {loss.item():.4f}")
losses.append(loss.item())
loss = loss / args.gradient_accumulation_steps
loss.backward()
'''
apart from the usual gradient accumulation steps,
we also do a backward pass after computing the last forward pass in the epoch (last_batch == True)
this is to make sure that we're not missing out on any data
'''
last_batch = (step + 1 == len(train_dataloader))
if (step + 1) % args.gradient_accumulation_steps == 0 or last_batch:
if optimizer is not None:
optimizer.step()
optimizer.zero_grad()
if optimizer_prod is not None:
optimizer_prod.step()
optimizer_prod.zero_grad()
# after every optimizer step, we reset the non-trainable embeddings to the original embeddings
embedding_handler.retract_embeddings(print_stds = (global_step % 50 == 0))
embedding_handler.fix_embedding_std(args.off_ratio_power)
# Track the learning rates for final plotting:
lora_lrs.append(get_avg_lr(optimizer_prod))
try:
ti_lrs.append(optimizer.param_groups[0]['lr'])
except:
ti_lrs.append(0.0)
# Print some statistics:
if (global_step % args.checkpointing_steps == 0): # and (global_step > 0):
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
save_lora(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=args.token_dict,
args_dict=args.args_dict,
is_lora= args.is_lora,
unet_lora_parameters=unet_lora_parameters,
unet_param_to_optimize_names=unet_param_to_optimize_names
)
args.save_as_json(os.path.join(output_save_dir,"training_args.json"))
last_save_step = global_step
validation_prompts = render_images(
pipe, target_size,
output_save_dir,
global_step,
args.seed,
args.is_lora,
args.pretrained_model,
n_imgs = 4
)
if args.debug:
token_embeddings = embedding_handler.get_trainable_embeddings()
for i, token_embeddings_i in enumerate(token_embeddings):
plot_torch_hist(
token_embeddings_i[0],
global_step,
args.output_dir,
f"embeddings_weights_token_0_{i}",
min_val=-0.05,
max_val=0.05,
ymax_f = 0.05
)
plot_torch_hist(
token_embeddings_i[1],
global_step,
args.output_dir,
f"embeddings_weights_token_1_{i}",
min_val=-0.05,
max_val=0.05,
ymax_f = 0.05
)
embedding_handler.print_token_info()
plot_torch_hist(
unet_lora_parameters,
global_step,
args.output_dir,
"lora_weights",
min_val=-0.3,
max_val=0.3,
ymax_f = 0.05
)
plot_loss(losses, save_path=f'{args.output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{args.output_dir}/learning_rates.png')
gc.collect()
torch.cuda.empty_cache()
images_done += args.train_batch_size
global_step += 1
if global_step % 100 == 0:
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
if args.debug:
plot_loss(losses, save_path=f'{args.output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{args.output_dir}/learning_rates.png')
plot_torch_hist(unet_lora_parameters, global_step, args.output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
plot_torch_hist(embedding_handler.get_trainable_embeddings(), global_step, args.output_dir, "embeddings_weights", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
# final_save
if (global_step - last_save_step) > 51:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
else:
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
if not os.path.exists(output_save_dir):
save_lora(
output_dir=output_save_dir,
global_step=global_step,
unet=unet,
embedding_handler=embedding_handler,
token_dict=args.token_dict,
args_dict=args.args_dict,
is_lora= args.is_lora,
unet_lora_parameters=unet_lora_parameters,
unet_param_to_optimize_names=unet_param_to_optimize_names
)
args.save_as_json(os.path.join(output_save_dir,"training_args.json"))
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, args.seed, args.is_lora, args.pretrained_model, n_imgs = 4, n_steps = 35)
else:
print(f"Skipping final save, {output_save_dir} already exists")
del unet
del vae
del text_encoder_one
del text_encoder_two
del tokenizer_one
del tokenizer_two
del embedding_handler
del pipe
gc.collect()
torch.cuda.empty_cache()
return output_save_dir
View File
+23
View File
@@ -0,0 +1,23 @@
def get_avg_lr(optimizer):
# Calculate the weighted average effective learning rate
total_lr = 0
total_params = 0
for group in optimizer.param_groups:
d = group['d']
lr = group['lr']
bias_correction = 1 # Default value
if group['use_bias_correction']:
beta1, beta2 = group['betas']
k = group['k']
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
effective_lr = d * lr * bias_correction
# Count the number of parameters in this group
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
total_lr += effective_lr * num_params
total_params += num_params
if total_params == 0:
return 0.0
else: return total_lr / total_params
+174
View File
@@ -0,0 +1,174 @@
import os, json
import torch
from safetensors.torch import load_file
from typing import Dict
from peft import PeftModel
from ..dataset_and_utils import TokenEmbeddingsHandler
from safetensors.torch import save_file
from .string import replace_in_string
'''
from diffusers.utils import (
convert_all_state_dict_to_peft,
convert_state_dict_to_diffusers,
convert_unet_state_dict_to_peft
)
'''
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
return training_args["trigger_text"] + ", " + prompt
lora_name_encapsulated = "<" + lora_name + ">"
trigger_text = training_args["trigger_text"]
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK"
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
# Fix common mistakes
fix_replacements = {
r",,": ",",
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
r"\s\.": ".",
r"\s,": ","
}
prompt = replace_in_string(prompt, fix_replacements)
if verbose:
print('-------------------------')
print("Adjusted prompt for LORA:")
print(orig_prompt)
print('-- to:')
print(prompt)
print('-------------------------')
return prompt
def patch_pipe_with_lora(pipe, lora_path):
"""
update the pipe with the lora model and the token embeddings
"""
pipe.unet = PeftModel.from_pretrained(pipe.unet, lora_path)
pipe.unet.merge_adapter()
# Load the textual_inversion token embeddings into the pipeline:
try: #SDXL
handler = TokenEmbeddingsHandler([pipe.text_encoder, pipe.text_encoder_2], [pipe.tokenizer, pipe.tokenizer_2])
except: #SD15
handler = TokenEmbeddingsHandler([pipe.text_encoder, None], [pipe.tokenizer, None])
embeddings_path = [f for f in os.listdir(lora_path) if f.endswith("embeddings.safetensors")][0]
handler.load_embeddings(os.path.join(lora_path, embeddings_path))
return pipe
def unet_attn_processors_state_dict(unet) -> Dict[str, torch.tensor]:
"""
Returns:
a state dict containing just the attention processor parameters.
"""
attn_processors = unet.attn_processors
attn_processors_state_dict = {}
for attn_processor_key, attn_processor in attn_processors.items():
for parameter_key, parameter in attn_processor.state_dict().items():
attn_processors_state_dict[
f"{attn_processor_key}.{parameter_key}"
] = parameter
return attn_processors_state_dict
def save_lora(output_dir, global_step, unet, embedding_handler, token_dict, args_dict, is_lora, unet_lora_parameters, unet_param_to_optimize_names):
"""
Save the LORA model to output_dir, optionally with some example images
"""
print(f"Saving checkpoint at step.. {global_step}")
os.makedirs(output_dir, exist_ok=True)
if not is_lora:
lora_tensors = {
name: param
for name, param in unet.named_parameters()
if name in unet_param_to_optimize_names
}
save_file(lora_tensors, f"{output_dir}/unet.safetensors",)
elif len(unet_lora_parameters) > 0:
unet.save_pretrained(save_directory = output_dir)
try:
concept_name = args_dict["name"].lower()
except:
concept_name = "eden_concept_lora"
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
concept_name = concept_name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
embedding_handler.save_embeddings(f"{output_dir}/{concept_name}_embeddings.safetensors",)
with open(f"{output_dir}/special_params.json", "w") as f:
json.dump(token_dict, f)
with open(f"{output_dir}/training_args.json", "w") as f:
json.dump(args_dict, f, indent=4)
+13
View File
@@ -0,0 +1,13 @@
def print_trainable_parameters(model, name = ''):
trainable_params = 0
all_param = 0
for _, param in model.named_parameters():
all_param += param.numel()
if param.requires_grad:
trainable_params += param.numel()
line_delimiter = "#" * 70
print('\n', line_delimiter)
print(
f"Trainable {name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%"
)
print(line_delimiter, '\n')
+120
View File
@@ -0,0 +1,120 @@
import random
import json
import os
import gc
import torch
from ..dataset_and_utils import load_models
from .lora import patch_pipe_with_lora, prepare_prompt_for_lora
from ..val_prompts import val_prompts
from diffusers import EulerDiscreteScheduler
from PIL import Image
def make_validation_img_grid(img_folder):
"""
find all the .jpg imgs in img_folder (template = *.jpg)
if >=4 validation imgs, create a 2x2 grid of them
otherwise just return the first validation img
"""
# Find all validation images
validation_imgs = sorted([f for f in os.listdir(img_folder) if f.endswith(".jpg")])
if len(validation_imgs) < 4:
# If less than 4 validation images, return path of the first one
return os.path.join(img_folder, validation_imgs[0])
else:
# If >= 4 validation images, create 2x2 grid
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:4]]
# Assuming all images are the same size, get dimensions of first image
width, height = imgs[0].size
# Create an empty image with 2x2 grid size
grid_img = Image.new("RGB", (2 * width, 2 * height))
# Paste the images into the grid
for i in range(2):
for j in range(2):
grid_img.paste(imgs.pop(0), (i * width, j * height))
# Save the new image
grid_img_path = os.path.join(img_folder, "validation_grid.jpg")
grid_img.save(grid_img_path)
return grid_img_path
@torch.no_grad()
def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0"):
random.seed(seed)
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
concept_mode = training_args["concept_mode"]
if concept_mode == "style":
validation_prompts_raw = random.sample(val_prompts['style'], n_imgs)
validation_prompts_raw[0] = ''
elif concept_mode == "face":
validation_prompts_raw = random.sample(val_prompts['face'], n_imgs)
validation_prompts_raw[0] = '<concept>'
else:
validation_prompts_raw = random.sample(val_prompts['object'], n_imgs)
validation_prompts_raw[0] = '<concept>'
reload_entire_pipeline = False
if reload_entire_pipeline: # reload the entire pipeline from disk and load in the lora module
print(f"Reloading entire pipeline from disk..")
gc.collect()
torch.cuda.empty_cache()
(pipeline,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet) = load_models(pretrained_model, device, torch.float16)
pipeline = pipeline.to(device)
pipeline = patch_pipe_with_lora(pipeline, lora_path)
else:
print(f"Re-using training pipeline for inference, just swapping the scheduler..")
pipeline = training_pipeline
training_scheduler = pipeline.scheduler
pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config)
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path) for prompt in validation_prompts_raw]
generator = torch.Generator(device=device).manual_seed(0)
pipeline_args = {
"negative_prompt": "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft",
"num_inference_steps": n_steps,
"guidance_scale": 7,
"height": render_size[0],
"width": render_size[1],
}
if is_lora > 0:
cross_attention_kwargs = {"scale": lora_scale}
else:
cross_attention_kwargs = None
for i in range(n_imgs):
pipeline_args["prompt"] = validation_prompts[i]
print(f"Rendering validation img with prompt: {validation_prompts[i]}")
image = pipeline(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
image.save(os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"), format="JPEG", quality=95)
# create img_grid:
img_grid_path = make_validation_img_grid(lora_path)
if not reload_entire_pipeline: # restore the training scheduler
pipeline.scheduler = training_scheduler
return validation_prompts_raw
+24
View File
@@ -0,0 +1,24 @@
def compute_snr(noise_scheduler, timesteps):
"""
Computes SNR as per
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
"""
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = alphas_cumprod**0.5
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
# Expand the tensors.
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
# Compute SNR.
snr = (alpha / sigma) ** 2
return snr
+14
View File
@@ -0,0 +1,14 @@
import re
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
+69 -782
View File
@@ -1,783 +1,70 @@
import fnmatch from trainer import TrainerConfig, Trainer
import json from preprocess import preprocess
import math
import os import os
import sys from io_utils import MODEL_DICT
import random
import time out_root_dir = "./lora_models"
import shutil run_name = "face_01"
import gc concept_mode = "face"
import numpy as np
from typing import List, Optional output_dir = os.path.join(out_root_dir, run_name)
import torch input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
import torch.utils.checkpoint output_dir,
import torch.nn.functional as F concept_mode = concept_mode,
input_zip_path = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
from peft import LoraConfig, get_peft_model #caption_text="in the style of TOK, ",
from diffusers.optimization import get_scheduler caption_text="",
from diffusers import EulerDiscreteScheduler mask_target_prompts=None,
from tqdm import tqdm target_size=1024,
crop_based_on_salience=True,
from dataset_and_utils import * use_face_detection_instead=False,
from lora_utils import * temp=0.7,
from io_utils import make_validation_img_grid left_right_flip_augmentation=False,
import matplotlib.pyplot as plt augment_imgs_up_to_n = 20,
seed = 0,
caption_model = "blip"
def print_trainable_parameters(model, name = ''): )
trainable_params = 0
all_param = 0 print('-------------------------------------------')
for _, param in model.named_parameters(): print(f"Trigger text: {trigger_text}")
all_param += param.numel() print(f'n_imgs: {n_imgs}')
if param.requires_grad: print(f'concept_mode: {concept_mode}')
trainable_params += param.numel() print('-------------------------------------------')
line_delimiter = "#" * 70
print('\n', line_delimiter)
print( config = TrainerConfig(
f"Trainable {name} params: {trainable_params/1000000:.1f}M || All params: {all_param/1000000:.1f}M || trainable = {100 * trainable_params / all_param:.2f}%" pretrained_model = MODEL_DICT['sdxl'],
) name='unnamed',
print(line_delimiter, '\n') concept_mode=concept_mode,
trigger_text=trigger_text,
instance_data_dir = os.path.join(input_dir, "captions.csv"),
def compute_snr(noise_scheduler, timesteps): output_dir = output_dir,
""" resolution= 1024,
Computes SNR as per train_batch_size = 4,
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849 max_train_steps = 600,
""" checkpointing_steps = 200,
alphas_cumprod = noise_scheduler.alphas_cumprod num_train_epochs = 10000,
sqrt_alphas_cumprod = alphas_cumprod**0.5 gradient_accumulation_steps = 1,
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5 textual_inversion_lr = 5e-4,
textual_inversion_weight_decay = 3e-4,
# Expand the tensors. lora_weight_decay = 0.00,
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026 prodigy_d_coef = 1.0,
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float() l1_penalty = 0.0,
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape): snr_gamma = 5.0,
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None] precision = "bf16",
alpha = sqrt_alphas_cumprod.expand(timesteps.shape) token_dict = {"TOK": "<s0><s1>"},
inserting_list_tokens = ["<s0>","<s1>"],
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float() is_lora = True,
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape): lora_rank = 12,
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None] lora_alpha = 12,
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape) hard_pivot = False,
off_ratio_power = 0.1,
# Compute SNR. args_dict = {},
snr = (alpha / sigma) ** 2 debug = True,
return snr seed = 0
)
def get_avg_lr(optimizer):
# Calculate the weighted average effective learning rate trainer = Trainer(config)
total_lr = 0 trainer.train()
total_params = 0 print("DONE")
for group in optimizer.param_groups:
d = group['d']
lr = group['lr']
bias_correction = 1 # Default value
if group['use_bias_correction']:
beta1, beta2 = group['betas']
k = group['k']
bias_correction = ((1 - beta2**(k+1))**0.5) / (1 - beta1**(k+1))
effective_lr = d * lr * bias_correction
# Count the number of parameters in this group
num_params = sum(p.numel() for p in group['params'] if p.requires_grad)
total_lr += effective_lr * num_params
total_params += num_params
if total_params == 0:
return 0.0
else: return total_lr / total_params
import re
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
return training_args["trigger_text"] + ", " + prompt
lora_name_encapsulated = "<" + lora_name + ">"
trigger_text = training_args["trigger_text"]
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK"
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
# Fix common mistakes
fix_replacements = {
r",,": ",",
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
r"\s\.": ".",
r"\s,": ","
}
prompt = replace_in_string(prompt, fix_replacements)
if verbose:
print('-------------------------')
print("Adjusted prompt for LORA:")
print(orig_prompt)
print('-- to:')
print(prompt)
print('-------------------------')
return prompt
from val_prompts import val_prompts
@torch.no_grad()
def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0"):
random.seed(seed)
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
concept_mode = training_args["concept_mode"]
if concept_mode == "style":
validation_prompts_raw = random.sample(val_prompts['style'], n_imgs)
validation_prompts_raw[0] = ''
elif concept_mode == "face":
validation_prompts_raw = random.sample(val_prompts['face'], n_imgs)
validation_prompts_raw[0] = '<concept>'
else:
validation_prompts_raw = random.sample(val_prompts['object'], n_imgs)
validation_prompts_raw[0] = '<concept>'
reload_entire_pipeline = False
if reload_entire_pipeline: # reload the entire pipeline from disk and load in the lora module
print(f"Reloading entire pipeline from disk..")
gc.collect()
torch.cuda.empty_cache()
(pipeline,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet) = load_models(pretrained_model, device, torch.float16)
pipeline = pipeline.to(device)
pipeline = patch_pipe_with_lora(pipeline, lora_path)
else:
print(f"Re-using training pipeline for inference, just swapping the scheduler..")
pipeline = training_pipeline
training_scheduler = pipeline.scheduler
pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config)
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path) for prompt in validation_prompts_raw]
generator = torch.Generator(device=device).manual_seed(0)
pipeline_args = {
"negative_prompt": "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft",
"num_inference_steps": n_steps,
"guidance_scale": 7,
"height": render_size[0],
"width": render_size[1],
}
if is_lora > 0:
cross_attention_kwargs = {"scale": lora_scale}
else:
cross_attention_kwargs = None
for i in range(n_imgs):
pipeline_args["prompt"] = validation_prompts[i]
print(f"Rendering validation img with prompt: {validation_prompts[i]}")
image = pipeline(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
image.save(os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"), format="JPEG", quality=95)
# create img_grid:
img_grid_path = make_validation_img_grid(lora_path)
if not reload_entire_pipeline: # restore the training scheduler
pipeline.scheduler = training_scheduler
return validation_prompts_raw
def main(
pretrained_model,
instance_data_dir: Optional[str] = "./dataset/zeke/captions.csv",
output_dir: str = "lora_output",
seed: Optional[int] = random.randint(0, 2**32 - 1),
resolution: int = 768,
crops_coords_top_left_h: int = 0,
crops_coords_top_left_w: int = 0,
train_batch_size: int = 1,
do_cache: bool = True,
num_train_epochs: int = 10000,
max_train_steps: Optional[int] = None,
checkpointing_steps: int = 500000, # default to no checkpoints
gradient_accumulation_steps: int = 1, # todo
unet_learning_rate: float = 1.0,
ti_lr: float = 3e-4,
lora_lr: float = 1.0,
prodigy_d_coef: float = 0.33,
l1_penalty: float = 0.0,
lora_weight_decay: float = 0.005,
ti_weight_decay: float = 0.001,
scale_lr: bool = False,
lr_scheduler: str = "constant",
lr_warmup_steps: int = 50,
lr_num_cycles: int = 1,
lr_power: float = 1.0,
snr_gamma: float = 5.0,
dataloader_num_workers: int = 0,
allow_tf32: bool = True,
mixed_precision: Optional[str] = "bf16",
device: str = "cuda:0",
token_dict: dict = {"TOKEN": "<s0>"},
inserting_list_tokens: List[str] = ["<s0>"],
verbose: bool = True,
is_lora: bool = True,
lora_rank: int = 8,
args_dict: dict = {},
debug: bool = False,
hard_pivot: bool = True,
off_ratio_power: float = 0.1,
) -> None:
if allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
print("Using seed", seed)
torch.manual_seed(seed)
weight_dtype = torch.float32
if mixed_precision == "fp16":
weight_dtype = torch.float16
elif mixed_precision == "bf16":
weight_dtype = torch.bfloat16
print(f"Loading models with weight_dtype: {weight_dtype}")
if scale_lr:
unet_learning_rate = (
unet_learning_rate * gradient_accumulation_steps * train_batch_size
)
(
pipe,
tokenizer_one,
tokenizer_two,
noise_scheduler,
text_encoder_one,
text_encoder_two,
vae,
unet,
) = load_models(pretrained_model, device, weight_dtype)
# Initialize new tokens for training.
embedding_handler = TokenEmbeddingsHandler(
[text_encoder_one, text_encoder_two], [tokenizer_one, tokenizer_two]
)
#starting_toks = ["person", "face"]
starting_toks = None
embedding_handler.initialize_new_tokens(inserting_toks=inserting_list_tokens, starting_toks=starting_toks, seed=seed)
text_encoders = [text_encoder_one, text_encoder_two]
unet_param_to_optimize = []
text_encoder_parameters = []
for text_encoder in text_encoders:
if text_encoder is not None:
for name, param in text_encoder.named_parameters():
if "token_embedding" in name:
param.requires_grad = True
text_encoder_parameters.append(param)
else:
param.requires_grad = False
unet_param_to_optimize_names = []
unet_lora_parameters = []
if not is_lora:
WHITELIST_PATTERNS = [
# "*.attn*.weight",
# "*ff*.weight",
"*"
]
BLACKLIST_PATTERNS = ["*.norm*.weight", "*time*"]
for name, param in unet.named_parameters():
if any(
fnmatch.fnmatch(name, pattern) for pattern in WHITELIST_PATTERNS
) and not any(
fnmatch.fnmatch(name, pattern) for pattern in BLACKLIST_PATTERNS
):
param.requires_grad_(True)
unet_param_to_optimize_names.append(name)
print(f"Training: {name}")
else:
param.requires_grad_(False)
# Optimizer creation
params_to_optimize = [
{
"params": text_encoder_parameters,
"lr": ti_lr,
"weight_decay": ti_weight_decay,
},
]
params_to_optimize_prodigy = [
{
"params": unet_param_to_optimize,
"lr": unet_learning_rate,
"weight_decay": lora_weight_decay,
},
]
else:
# Do lora-training instead.
unet.requires_grad_(False)
# https://huggingface.co/docs/peft/main/en/developer_guides/lora#rank-stabilized-lora
unet_lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_rank,
init_lora_weights="gaussian",
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
#use_rslora=True,
use_dora=True,
)
#unet.add_adapter(unet_lora_config)
unet = get_peft_model(unet, unet_lora_config)
print_trainable_parameters(unet, name = 'unet')
unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters()))
params_to_optimize = [
{
"params": text_encoder_parameters,
"lr": ti_lr,
"weight_decay": ti_weight_decay,
},
]
params_to_optimize_prodigy = [
{
"params": unet_lora_parameters,
"lr": 1.0,
"weight_decay": lora_weight_decay,
},
]
optimizer_type = "prodigy" # hardcode for now
if optimizer_type != "prodigy":
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=0.0, # this wd doesn't matter, I think
)
else:
try:
import prodigyopt
except ImportError:
raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`")
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer_prod = prodigyopt.Prodigy(
params_to_optimize_prodigy,
d_coef = prodigy_d_coef,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=lora_weight_decay,
betas=(0.9, 0.99),
growth_rate=1.025, # this slows down the lr_rampup
)
optimizer = torch.optim.AdamW(
params_to_optimize,
weight_decay=ti_weight_decay,
)
train_dataset = PreprocessedDataset(
instance_data_dir,
tokenizer_one,
tokenizer_two,
vae,
do_cache=True,
substitute_caption_map=token_dict,
)
print(f"# PTI : Loaded dataset, do_cache: {do_cache}")
train_dataloader = torch.utils.data.DataLoader(
train_dataset,
batch_size=train_batch_size,
shuffle=True,
num_workers=dataloader_num_workers,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / gradient_accumulation_steps
)
if max_train_steps is None:
max_train_steps = num_train_epochs * num_update_steps_per_epoch
lr_scheduler = get_scheduler(
lr_scheduler,
optimizer=optimizer,
num_warmup_steps=lr_warmup_steps * gradient_accumulation_steps,
num_training_steps=max_train_steps * gradient_accumulation_steps,
num_cycles=lr_num_cycles,
power=lr_power,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / gradient_accumulation_steps
)
num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch)
total_batch_size = train_batch_size * gradient_accumulation_steps
if verbose:
print(f"# PTI : Running training ")
print(f"# PTI : Num examples = {len(train_dataset)}")
print(f"# PTI : Num batches each epoch = {len(train_dataloader)}")
print(f"# PTI : Num Epochs = {num_train_epochs}")
print(f"# PTI : Instantaneous batch size per device = {train_batch_size}")
print(
f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}"
)
print(f"# PTI : Gradient Accumulation steps = {gradient_accumulation_steps}")
print(f"# PTI : Total optimization steps = {max_train_steps}")
global_step = 0
first_epoch = 0
last_save_step = 0
progress_bar = tqdm(range(global_step, max_train_steps), position=0, leave=True)
checkpoint_dir = os.path.join(str(output_dir), "checkpoints")
if os.path.exists(checkpoint_dir):
shutil.rmtree(checkpoint_dir)
os.makedirs(f"{checkpoint_dir}")
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
#embedding_handler.pre_optimize_token_embeddings(train_dataset)
ti_lrs, lora_lrs = [], []
losses = []
start_time, images_done = time.time(), 0
for epoch in range(first_epoch, num_train_epochs):
unet.train()
progress_bar.set_description(f"# PTI :step: {global_step}, epoch: {epoch}")
for step, batch in enumerate(train_dataloader):
progress_bar.update(1)
if hard_pivot:
if epoch >= num_train_epochs // 2:
if optimizer is not None:
print("----------------------")
print("# PTI : Pivot halfway")
print("----------------------")
# remove text encoder parameters from the optimizer
optimizer.param_groups = None
# remove the optimizer state corresponding to text_encoder_parameters
for param in text_encoder_parameters:
if param in optimizer.state:
del optimizer.state[param]
optimizer = None
else: # Update learning rates gradually:
finegrained_epoch = epoch + step / len(train_dataloader)
completion_f = finegrained_epoch / num_train_epochs
# param_groups[1] goes from ti_lr to 0.0 over the course of training
optimizer.param_groups[0]['lr'] = ti_lr * (1 - completion_f) ** 2.0
try: #sdxl
(tok1, tok2), vae_latent, mask = batch
except: #sd15
tok1, vae_latent, mask = batch
tok2 = None
vae_latent = vae_latent.to(weight_dtype)
# tokens to text embeds
prompt_embeds_list = []
for tok, text_encoder in zip((tok1, tok2), text_encoders):
if tok is None:
continue
prompt_embeds_out = text_encoder(
tok.to(text_encoder.device),
output_hidden_states=True,
)
pooled_prompt_embeds = prompt_embeds_out[0]
prompt_embeds = prompt_embeds_out.hidden_states[-2]
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(bs_embed, seq_len, -1)
prompt_embeds_list.append(prompt_embeds)
prompt_embeds = torch.concat(prompt_embeds_list, dim=-1)
pooled_prompt_embeds = pooled_prompt_embeds.view(bs_embed, -1)
# Create Spatial-dimensional conditions.
original_size = (resolution, resolution)
target_size = (resolution, resolution)
crops_coords_top_left = (crops_coords_top_left_h, crops_coords_top_left_w)
add_time_ids = list(original_size + crops_coords_top_left + target_size)
add_time_ids = torch.tensor([add_time_ids])
add_time_ids = add_time_ids.to(device, dtype=prompt_embeds.dtype).repeat(
bs_embed, 1
)
# Sample noise that we'll add to the latents:
noise = torch.randn_like(vae_latent)
noise_offset = 0.05 # TODO, turn this into an input arg and do a grid search
if noise_offset > 0.0:
# https://www.crosslabs.org//blog/diffusion-with-offset-noise
noise += noise_offset * torch.randn(
(noise.shape[0], noise.shape[1], 1, 1), device=noise.device)
bsz = vae_latent.shape[0]
timesteps = torch.randint(
0,
noise_scheduler.config.num_train_timesteps,
(bsz,),
device=vae_latent.device,
).long()
noisy_model_input = noise_scheduler.add_noise(vae_latent, noise, timesteps)
noise_sigma = 0.0
if noise_sigma > 0.0: # experimental: apply random noise to the conditioning vectors as a form of regularization
prompt_embeds[0,1:-2,:] += torch.randn_like(prompt_embeds[0,1:-2,:]) * noise_sigma
# Predict the noise residual
model_pred = unet(
noisy_model_input,
timesteps,
prompt_embeds,
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
).sample
# Get the unet prediction target depending on the prediction type:
if noise_scheduler.config.prediction_type == "epsilon":
target = noise
elif noise_scheduler.config.prediction_type == "v_prediction":
target = noise_scheduler.get_velocity(model_input, noise, timesteps)
else:
raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}")
# Compute the loss:
if snr_gamma is None:
loss = (model_pred - target).pow(2) * mask
# modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
# Average the normalized errors across the batch
loss = loss.mean()
else:
# Compute loss-weights as per Section 3.4 of https://arxiv.org/abs/2303.09556.
# Since we predict the noise instead of x_0, the original formulation is slightly changed.
# This is discussed in Section 4.2 of the same paper.
snr = compute_snr(noise_scheduler, timesteps)
base_weight = (
torch.stack([snr, snr_gamma * torch.ones_like(timesteps)], dim=1).min(dim=1)[0] / snr
)
if noise_scheduler.config.prediction_type == "v_prediction":
# Velocity objective needs to be floored to an SNR weight of one.
mse_loss_weights = base_weight + 1
else:
# Epsilon and sample both use the same loss weights.
mse_loss_weights = base_weight
mse_loss_weights = mse_loss_weights / mse_loss_weights.mean()
loss = (model_pred - target).pow(2) * mask
loss = loss.mean(dim=list(range(1, len(loss.shape)))) * mse_loss_weights
if 1: # modulate loss by the inverse of the mask's mean value
mean_mask_values = mask.mean(dim=list(range(1, len(loss.shape))))
mean_mask_values = mean_mask_values / mean_mask_values.mean()
loss = loss.mean(dim=list(range(1, len(loss.shape)))) / mean_mask_values
loss = loss.mean()
if 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 += l1_penalty * l1_norm
losses.append(loss.item())
loss = loss / gradient_accumulation_steps
loss.backward()
'''
apart from the usual gradient accumulation steps,
we also do a backward pass after computing the last forward pass in the epoch (last_batch == True)
this is to make sure that we're not missing out on any data
'''
last_batch = (step + 1 == len(train_dataloader))
if (step + 1) % gradient_accumulation_steps == 0 or last_batch:
if optimizer is not None:
optimizer.step()
optimizer.zero_grad()
optimizer_prod.step()
optimizer_prod.zero_grad()
# after every optimizer step, we reset the non-trainable embeddings to the original embeddings
embedding_handler.retract_embeddings(print_stds = (global_step % 50 == 0))
embedding_handler.fix_embedding_std(off_ratio_power)
# Track the learning rates for final plotting:
lora_lrs.append(get_avg_lr(optimizer_prod))
try:
ti_lrs.append(optimizer.param_groups[0]['lr'])
except:
ti_lrs.append(0.0)
# Print some statistics:
if (global_step % checkpointing_steps == 0):
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
save_lora(output_save_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names)
last_save_step = global_step
if debug:
token_embeddings = embedding_handler.get_trainable_embeddings()
for i, token_embeddings_i in enumerate(token_embeddings):
plot_torch_hist(token_embeddings_i[0], global_step, output_dir, f"embeddings_weights_token_0_{i}", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
plot_torch_hist(token_embeddings_i[1], global_step, output_dir, f"embeddings_weights_token_1_{i}", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
embedding_handler.print_token_info()
plot_torch_hist(unet_lora_parameters, global_step, output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
plot_loss(losses, save_path=f'{output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{output_dir}/learning_rates.png')
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, seed, is_lora, pretrained_model, n_imgs = 4)
gc.collect()
torch.cuda.empty_cache()
images_done += train_batch_size
global_step += 1
if global_step % 100 == 0:
print(f" ---- avg training fps: {images_done / (time.time() - start_time):.2f}", end="\r")
if global_step % (max_train_steps//20) == 0:
progress = (global_step / max_train_steps) + 0.05
yield np.min((progress, 1.0))
# final_save
if (global_step - last_save_step) > 51:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
else:
output_save_dir = f"{checkpoint_dir}/checkpoint-{last_save_step}"
if debug:
plot_loss(losses, save_path=f'{output_dir}/losses.png')
plot_lrs(lora_lrs, ti_lrs, save_path=f'{output_dir}/learning_rates.png')
plot_torch_hist(unet_lora_parameters, global_step, output_dir, "lora_weights", min_val=-0.3, max_val=0.3, ymax_f = 0.05)
plot_torch_hist(embedding_handler.get_trainable_embeddings(), global_step, output_dir, "embeddings_weights", min_val=-0.05, max_val=0.05, ymax_f = 0.05)
if not os.path.exists(output_save_dir):
save_lora(output_save_dir, global_step, unet, embedding_handler, token_dict, args_dict, seed, is_lora, unet_lora_parameters, unet_param_to_optimize_names)
validation_prompts = render_images(pipe, target_size, output_save_dir, global_step, seed, is_lora, pretrained_model, n_imgs = 4, n_steps = 35)
else:
print(f"Skipping final save, {output_save_dir} already exists")
del unet
del vae
del text_encoder_one
del text_encoder_two
del tokenizer_one
del tokenizer_two
del embedding_handler
del pipe
gc.collect()
torch.cuda.empty_cache()
with open(f"{output_save_dir}/training_args.json", "w") as f:
args_dict["grid_prompts"] = validation_prompts
json.dump(args_dict, f, indent=4)
return output_save_dir, validation_prompts
if __name__ == "__main__":
main()