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
.huggingface
tests/
trainer/
train.py
debug/*
!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
this random variance, eg CLIP_similarity pretraining.
- Improve the img captioning by swapping BLIP for cogVLM: https://github.com/THUDM/CogVLM
- it looks the like gpu-utilization is only like 65-70% during training: whats the bottleneck? Can we speed this up?
Bugfixing:
see msgs at: https://discord.com/channels/573691888050241543/1184175211998883950/1217550596878373037
@@ -50,6 +51,7 @@ Bigger improvements:
Tuning Experiments once code is fully ready:
- test if VAE weight_type actually matters for training
- try-out conditioning noise injection during training to increase robustness
- re-test / tweak the adaptive learning rates instead of hard-pivot (also test Prodigy vs Adam)
- right now it looks like the diffusion model gets partially "destroyed" in the beginning of training (outputs from steps 100-200 look terrible),
+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"
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"
SDXL_TURBO_MODEL_CACHE = "./models/SDXL_turbo.safetensors"
SDXL_TURBO_URL = "https://huggingface.co/stabilityai/sdxl-turbo/resolve/main/sd_xl_turbo_1.0_fp16.safetensors?download=true"
SDXL_LIGHTNING_MODEL_CACHE = "./models/SDXL_lightning.safetensors"
SDXL_LIGHTNING_URL = "https://huggingface.co/ByteDance/SDXL-Lightning/resolve/main/sdxl_lightning_8step.safetensors?download=true"
# Define model paths and URLs in a dictionary
MODEL_INFO = {
"sdxl": {"path": SDXL_MODEL_CACHE, "url": SDXL_URL},
"sd15": {"path": SD15_MODEL_CACHE, "url": SD15_URL}
MODEL_DICT = {
"sdxl": {
"path": SDXL_MODEL_CACHE,
"url": SDXL_URL,
"version": "sdxl"
},
"sd15": {
"path": SD15_MODEL_CACHE,
"url": SD15_URL,
"version": "sd15"
},
"sdxl_turbo": {
"path": SDXL_TURBO_MODEL_CACHE,
"url": SDXL_TURBO_URL,
"version": "sdxl_turbo"
},
"sdxl_lightning": {
"path": SDXL_LIGHTNING_MODEL_CACHE,
"url": SDXL_LIGHTNING_URL,
"version": "sdxl_lightning"
}
}
def download_weights(url, dest):
-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.xlabel('Weight Value')
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.close()
@@ -344,7 +344,7 @@ class TokenEmbeddingsHandler:
if text_encoder is None:
continue
trainable_embeddings.append(text_encoder.text_model.embeddings.token_embedding.weight.data[self.train_ids])
return trainable_embeddings
def find_nearest_tokens(self, query_embedding, tokenizer, text_encoder, idx, distance_metric, top_k = 5):
@@ -486,7 +486,6 @@ class TokenEmbeddingsHandler:
print("Initializing new tokens...")
print(inserting_toks)
torch.manual_seed(seed)
idx = 0
for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders):
@@ -541,6 +540,7 @@ class TokenEmbeddingsHandler:
else:
if 1: # random initialization:
torch.manual_seed(seed)
init_embeddings = (torch.randn(len(self.train_ids), text_encoder.text_model.config.hidden_size).to(device=self.device).to(dtype=self.dtype) * std_token_embedding * 1.0)
else:
# Test code to initialize the new tokens with some specific tokens
+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
import json
import math
from trainer import TrainerConfig, Trainer
from preprocess import preprocess
import os
import sys
import random
import time
import shutil
import gc
import numpy as np
from typing import List, Optional
import torch
import torch.utils.checkpoint
import torch.nn.functional as F
from peft import LoraConfig, get_peft_model
from diffusers.optimization import get_scheduler
from diffusers import EulerDiscreteScheduler
from tqdm import tqdm
from dataset_and_utils import *
from lora_utils import *
from io_utils import make_validation_img_grid
import matplotlib.pyplot as plt
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')
def compute_snr(noise_scheduler, timesteps):
"""
Computes SNR as per
https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
"""
alphas_cumprod = noise_scheduler.alphas_cumprod
sqrt_alphas_cumprod = alphas_cumprod**0.5
sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5
# Expand the tensors.
# Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026
sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None]
alpha = sqrt_alphas_cumprod.expand(timesteps.shape)
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float()
while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape):
sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None]
sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape)
# Compute SNR.
snr = (alpha / sigma) ** 2
return snr
def 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
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()
from io_utils import MODEL_DICT
out_root_dir = "./lora_models"
run_name = "face_01"
concept_mode = "face"
output_dir = os.path.join(out_root_dir, run_name)
input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
output_dir,
concept_mode = concept_mode,
input_zip_path = "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip",
#caption_text="in the style of TOK, ",
caption_text="",
mask_target_prompts=None,
target_size=1024,
crop_based_on_salience=True,
use_face_detection_instead=False,
temp=0.7,
left_right_flip_augmentation=False,
augment_imgs_up_to_n = 20,
seed = 0,
caption_model = "blip"
)
print('-------------------------------------------')
print(f"Trigger text: {trigger_text}")
print(f'n_imgs: {n_imgs}')
print(f'concept_mode: {concept_mode}')
print('-------------------------------------------')
config = TrainerConfig(
pretrained_model = MODEL_DICT['sdxl'],
name='unnamed',
concept_mode=concept_mode,
trigger_text=trigger_text,
instance_data_dir = os.path.join(input_dir, "captions.csv"),
output_dir = output_dir,
resolution= 1024,
train_batch_size = 4,
max_train_steps = 600,
checkpointing_steps = 200,
num_train_epochs = 10000,
gradient_accumulation_steps = 1,
textual_inversion_lr = 5e-4,
textual_inversion_weight_decay = 3e-4,
lora_weight_decay = 0.00,
prodigy_d_coef = 1.0,
l1_penalty = 0.0,
snr_gamma = 5.0,
precision = "bf16",
token_dict = {"TOK": "<s0><s1>"},
inserting_list_tokens = ["<s0>","<s1>"],
is_lora = True,
lora_rank = 12,
lora_alpha = 12,
hard_pivot = False,
off_ratio_power = 0.1,
args_dict = {},
debug = True,
seed = 0
)
trainer = Trainer(config)
trainer.train()
print("DONE")