diff --git a/main_sd3.py b/main_sd3.py index df0b940..65b3237 100644 --- a/main_sd3.py +++ b/main_sd3.py @@ -41,8 +41,6 @@ from transformers import ( from diffusers import ( StableDiffusion3Pipeline, FlowMatchEulerDiscreteScheduler, - SD3Transformer2DModel, - AutoencoderKL ) ## shared components with the other trainer (sdxl/sd15) @@ -50,13 +48,6 @@ from trainer.embedding_handler import TokenEmbeddingsHandler from trainer.loss import ( ConditioningRegularizer ) -from trainer.optimizer import ( - OptimizerCollection, - get_optimizer_and_peft_models_text_encoder_lora, - get_textual_inversion_optimizer, - get_unet_lora_parameters, - get_unet_optimizer -) from trainer.optimizer import count_trainable_params from peft import LoraConfig, get_peft_model from typing import Iterable @@ -70,7 +61,7 @@ from tqdm import tqdm import wandb from peft.utils import get_peft_model_state_dict -TRAIN_TRANSFORMER = False +TRAIN_TRANSFORMER = True TRAIN_TEXTUAL_INVERSION = True def load_sd3_tokenizers(): @@ -350,6 +341,7 @@ class TextualInversion: num_tokens: int = 2, device = "cuda:0", embedding_size: int = None, + initial_embed_string = None ): self.embedding_module = embedding_module self.tokenizer = tokenizer @@ -358,8 +350,6 @@ class TextualInversion: self.device = device self.embedding_size = embedding_size assert self.num_tokens>0 - print("embedding_module", embedding_module) - print(f"tokenizer", tokenizer) if embedding_size is not None: assert embedding_size>0 self.embedding_size = embedding_size @@ -370,7 +360,17 @@ class TextualInversion: self.embedding_size = foo.shape[-1] print(f"Auto-determined embedding_size to be: {self.embedding_size}") - embed_tensor = torch.randn(self.num_tokens, self.embedding_size).to(self.device) + if initial_embed_string is None: + embed_tensor = torch.randn(self.num_tokens, self.embedding_size).to(self.device) + else: + """ + The starting point of the TI token is the embedding corresponding to initial_embed_string + """ + # [0,1:-1,:] -> 0 means: remove the first (batch) dim, 1:-1 means skip the START and END tokens + with torch.no_grad(): + embed_tensor = self.tokenize_and_embed(text = initial_embed_string, padded = False)[0,1:-1,:] + + assert embed_tensor.shape[0] == num_tokens, f"Expected embed_tensor to have {num_tokens} tokens but got: {embed_tensor.shape[0]}. Try changing the num_tokens arg to {embed_tensor.shape[0]} to ignore this error." embed_tensor.requires_grad = True self.params = torch.nn.Parameter(embed_tensor) @@ -504,18 +504,20 @@ def main(config: TrainingConfig, wandb_log = False): embedding_module=pipeline.text_encoder.text_model.embeddings, tokenizer=pipeline.tokenizer, trigger_text = ", ", - num_tokens = 2 + num_tokens = 2, + initial_embed_string = "Belgian Man" ) textual_inversion_2 = TextualInversion( embedding_module=pipeline.text_encoder_2.text_model.embeddings, tokenizer=pipeline.tokenizer_2, trigger_text = ", ", - num_tokens = 2 + num_tokens = 2, + initial_embed_string = "Belgian Man" ) textual_inversion_params = [textual_inversion.params, textual_inversion_2.params] optimizer_textual_inversion = torch.optim.SGD( textual_inversion_params, - lr = 1e-3, + lr = 1e-4, weight_decay = 0.0 ) # embeds = textual_inversion.compute_text_embeddings(text = "ATOKA") @@ -548,6 +550,7 @@ def main(config: TrainingConfig, wandb_log = False): """ override some config params because we're recycling an sdxl config here """ + config.ti_lr_warmup_steps = 200 config.is_lora = True config.token_warmup_steps = 0 config.lora_rank = 8 @@ -580,8 +583,8 @@ def main(config: TrainingConfig, wandb_log = False): ) inference_prompts = [ - # ", A man is eating popcorn while holding a knife", - # ", A man is taking a selfie in space", + ", A man is eating popcorn while holding a knife", + ", A man is taking a selfie in space", ", Gentleman with a moustache dressed up as santa", ] @@ -687,14 +690,13 @@ def main(config: TrainingConfig, wandb_log = False): Scale learning rate of textual inversion params """ if config.ti_optimizer != "prodigy": # Update ti_learning rate gradually: - # optimizer_ti.param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0 + optimizer_textual_inversion.param_groups[0]['lr'] = config.ti_lr * (1 - completion_f) ** 2.0 # warmup the ti-lr: if config.ti_lr_warmup_steps > 0: warmup_f = min(global_step / config.ti_lr_warmup_steps, 1.0) - # optimizer_ti.param_groups[0]['lr'] *= warmup_f + optimizer_textual_inversion.param_groups[0]['lr'] *= warmup_f if config.freeze_ti_after_completion_f <= completion_f: - pass - # optimizer_ti.param_groups[0]['lr'] *= 0 + optimizer_textual_inversion.param_groups[0]['lr'] *= 0 if not config.aspect_ratio_bucketing: captions, vae_latent, mask = batch @@ -817,11 +819,11 @@ def main(config: TrainingConfig, wandb_log = False): torch.nn.utils.clip_grad_norm_(transformer_trainable_params, max_norm=1) if TRAIN_TEXTUAL_INVERSION: - print( - f"Textual Inversion grad norms: {textual_inversion.params.grad.norm()} | {textual_inversion_2.params.grad.norm()}" - ) + torch.nn.utils.clip_grad_norm_([textual_inversion.params, textual_inversion_2.params], max_norm=1) - optimizer.step() + if global_step % config.gradient_accumulation_steps == 0: + optimizer.step() + print(f"Performed an optimization step") progress_bar.set_postfix( { @@ -841,10 +843,16 @@ def main(config: TrainingConfig, wandb_log = False): "global_step": global_step, } # data["textual_inversion_lr"] = optimizer_ti.param_groups[0]['lr'] - data["transformer_lr"] = optimizer_transformer.param_groups[0]['lr'] - data["transformer_grad_norms"] = wandb.Histogram( - transformer_grad_norms, - ) + if TRAIN_TEXTUAL_INVERSION: + data["textual_inversion_lr"] = optimizer_textual_inversion.param_groups[0]['lr'] + data["textual_inversion_grad_norm"] = textual_inversion.params.grad.norm() + data["textual_inversion_2_grad_norm"] = textual_inversion_2.params.grad.norm() + + if TRAIN_TRANSFORMER: + data["transformer_lr"] = optimizer_transformer.param_groups[0]['lr'] + data["transformer_grad_norms"] = wandb.Histogram( + transformer_grad_norms, + ) wandb.log( data