diff --git a/main.py b/main.py index 0f5d9db..6c7b4ca 100755 --- a/main.py +++ b/main.py @@ -1,4 +1,3 @@ -import fnmatch import math import os import time @@ -12,8 +11,6 @@ import torch import torch.utils.checkpoint from tqdm import tqdm -from typing import Union, Iterable, List, Dict, Tuple, Optional, cast - from trainer.utils.utils import * from trainer.checkpoint import save_checkpoint from trainer.embedding_handler import TokenEmbeddingsHandler @@ -175,7 +172,10 @@ def train(config: TrainingConfig): aspect_ratio_bucketing=config.aspect_ratio_bucketing, train_batch_size=config.train_batch_size ) - # offload the vae to cpu: + print("Final training captions:") + print(train_dataset.captions[:40]) + + # offload the vae to cpu and release memory: vae = vae.to('cpu') gc.collect() torch.cuda.empty_cache() @@ -185,16 +185,10 @@ def train(config: TrainingConfig): train_dataset, batch_size=config.train_batch_size, shuffle=True, - num_workers=config.dataloader_num_workers, + num_workers=config.dataloader_num_workers ) - num_update_steps_per_epoch = math.ceil(len(train_dataloader) / config.gradient_accumulation_steps) - num_update_steps_per_epoch = math.ceil(len(train_dataloader)) - - if config.max_train_steps is None: - config.max_train_steps = config.num_train_epochs * num_update_steps_per_epoch - - config.num_train_epochs = math.ceil(config.max_train_steps / num_update_steps_per_epoch) + config.num_train_epochs = int(math.ceil(config.max_train_steps / len(train_dataloader))) total_batch_size = config.train_batch_size * config.gradient_accumulation_steps print(f"--- Num samples = {len(train_dataset)}") @@ -288,6 +282,8 @@ def train(config: TrainingConfig): else: captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch() + mask = mask.to(config.device) + captions = list(captions) prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals( config, pipe, captions @@ -327,11 +323,7 @@ def train(config: TrainingConfig): token_attention_loss = compute_token_attention_loss(pipe, embedding_handler, captions, mask, daam_loss) losses['token_attention_loss'].append(token_attention_loss.item()) - loss = loss + 0.0000002 * token_attention_loss - - if global_step%40 == 0: - img_ratio = config.train_img_size[0] / config.train_img_size[1] - plot_token_attention_loss(config.output_dir, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio) + loss = loss + config.token_attention_loss_w * token_attention_loss if config.training_attributes["gpt_description"] and config.debug: concept_description_loss = embedding_handler.compute_target_prompt_loss(config.training_attributes["gpt_description"], prompt_embeds, pooled_prompt_embeds, config, pipe) @@ -381,6 +373,10 @@ def train(config: TrainingConfig): for std_i, std in enumerate(embedding_stds): token_stds[f'text_encoder_{idx}'][std_i].append(embedding_stds[std_i].item()) + if global_step % 50 == 0: + img_ratio = config.train_img_size[0] / config.train_img_size[1] + plot_token_attention_loss(config.output_dir, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio) + # Print some statistics: if (global_step % config.checkpointing_steps == 0) and (global_step < (config.max_train_steps - 25)) and global_step > 0: @@ -404,6 +400,7 @@ def train(config: TrainingConfig): last_save_step = global_step if config.debug: + token_embeddings, trainable_tokens = embedding_handler.get_trainable_embeddings() for idx, text_encoder in enumerate(text_encoders): if text_encoder is None: @@ -425,7 +422,8 @@ def train(config: TrainingConfig): plot_grad_norms(grad_norms, save_path=f'{config.output_dir}/grad_norms.png') plot_lrs(optimizer_collection.learning_rate_tracker, save_path=f'{config.output_dir}/learning_rates.png') plot_curve(prompt_embeds_norms, 'steps', 'norm', 'prompt_embed norms', save_path=f'{config.output_dir}/prompt_embeds_norms.png') - + + validation_prompts = render_images( pipe = pipe, render_size = config.validation_img_size, diff --git a/scripts/create_hyperparam_sweep.py b/scripts/create_hyperparam_sweep.py index 1313a01..68155d6 100644 --- a/scripts/create_hyperparam_sweep.py +++ b/scripts/create_hyperparam_sweep.py @@ -35,12 +35,12 @@ def hamming_distance(dict1, dict2): ####################################################################################### # Setup the base experiment config: -exp_name = "beeple" +exp_name = "objects" caption_prefix = "" mask_target_prompts = "" n_exp = 200 # how many random experiment settings to generate -min_hamming_distance = 1 # min_n_params that have to be different from any previous experiment to be scheduled -nohup = True +min_hamming_distance = 2 # min_n_params that have to be different from any previous experiment to be scheduled +nohup = False output_sh_path = f"gridsearch_configs/{exp_name}.sh" # Define training hyperparameters and their possible values @@ -50,41 +50,45 @@ hyperparameters = { "output_dir": [f"lora_models/{exp_name}"], "sd_model_version": ["sdxl"], "lora_training_urls": [ - "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx.zip", + "/home/rednax/Documents/datasets/plantoid/plantoid", + "/home/rednax/Documents/datasets/sweep/banny", + "/home/rednax/Documents/datasets/sweep/banny_mini" + ], - "concept_mode": ['style'], + "concept_mode": ['object'], "sample_imgs_lora_scale": [0.8], - "disable_ti": ['false', 'true'], + "disable_ti": ['false'], "seed": [0], "resolution": [512], "train_batch_size": [4], - "n_sample_imgs": [8], - "max_train_steps": [1200], - "checkpointing_steps": [200], + "n_sample_imgs": [6], + "max_train_steps": [400], + "checkpointing_steps": [100], "gradient_accumulation_steps": [1], - "n_tokens": [2], + "n_tokens": [2,3,4], "ti_lr": [0.001], - "ti_weight_decay": [0.001], + "ti_weight_decay": [0.000], "l1_penalty": [0.0], "token_warmup_steps": [0], - "tok_cov_reg_w": [2000], + "tok_cov_reg_w": [500], + "token_attention_loss_w": [0, 2e-7, 10e-7], - "unet_lr": [0.0002, 0.00005], + "unet_lr": [0.001, 0.0003, 0.0001], "lora_alpha_multiplier": [1.0], "prodigy_d_coef": [1.0], "lora_weight_decay": [0.001], - "lora_rank": [16], + "lora_rank": [16,32], "use_dora": ['false'], - "unet_optimizer_type": ['AdamW8bit'], - "is_lora": ['false'], + "unet_optimizer_type": ['adamw'], + "is_lora": ['true'], "text_encoder_lora_optimizer": [None], "text_encoder_lora_lr": [0.0e-4], "snr_gamma": [5.0], - "caption_model": ["blip", "gpt4-v"], + "caption_model": ["blip", "florence"], "augment_imgs_up_to_n": [40], "verbose": ['true'], "debug": ['true'] diff --git a/scripts/parse_results.py b/scripts/parse_results.py new file mode 100644 index 0000000..3fa498d --- /dev/null +++ b/scripts/parse_results.py @@ -0,0 +1,146 @@ +import os +import json +import matplotlib.pyplot as plt +from collections import defaultdict + +# Step 1: Import necessary libraries (done above) + +# Step 2: Define a function to count JPG files in a directory +def count_jpg_files(directory): + return len([f for f in os.listdir(directory) if f.lower().endswith('.jpg')]) + +# Step 3: Define a function to find and load the training_args.json file +def load_training_args(directory): + for root, dirs, files in os.walk(directory): + if 'training_args.json' in files: + with open(os.path.join(root, 'training_args.json'), 'r') as f: + return json.load(f) + return None + +# Step 4: Traverse the directory structure and collect data +def collect_data(root_dir): + data = [] + for root, dirs, files in os.walk(root_dir): + if 'checkpoints' in dirs: + checkpoints_dir = os.path.join(root, 'checkpoints') + score = count_jpg_files(checkpoints_dir) + training_args = load_training_args(root) + if training_args: + data.append((training_args, score)) + else: + print(f"Warning: No training_args.json found in {root}") + print(f"Collected data from {len(data)} runs") + return data + +# Step 5: Process the collected data to identify varying hyperparameters +def identify_varying_hyperparams(data): + all_params = set().union(*[set(args.keys()) for args, _ in data]) + varying_params = {} + + def make_hashable(val): + if isinstance(val, dict): + return tuple(sorted((k, make_hashable(v)) for k, v in val.items())) + elif isinstance(val, list): + return tuple(make_hashable(v) for v in val) + elif isinstance(val, set): + return frozenset(make_hashable(v) for v in val) + return val + + for param in all_params: + try: + values = set(make_hashable(args.get(param)) for args, _ in data if param in args) + if len(values) > 1: + varying_params[param] = values + print(f"---> Parameter '{param}' varies across runs") + except TypeError as e: + print(f"Warning: Could not hash values for parameter '{param}'. Error: {e}") + print(f"Values: {[args.get(param) for args, _ in data if param in args]}") + + return varying_params + + +# Step 6: Create visual plots for each varying hyperparameter +import numpy as np +from scipy import stats +import matplotlib.pyplot as plt +from collections import defaultdict + +def create_plots(data, varying_params): + for param, values in varying_params.items(): + plt.figure(figsize=(12, 8)) + param_data = defaultdict(list) + + for args, score in data: + if param in args: + value = args[param] + value_str = str(value) + param_data[value_str].append(score) + + values_list = sorted(param_data.keys()) + all_x = [] + all_y = [] + + for i, value_str in enumerate(values_list): + scores = param_data[value_str] + jittered_x = np.random.normal(i, 0.1, size=len(scores)) + plt.scatter(jittered_x, scores, alpha=0.6, label=value_str) + all_x.extend([i] * len(scores)) + all_y.extend(scores) + + # Calculate trendline + x = np.array(all_x) + y = np.array(all_y) + z = np.polyfit(x, y, 1) + p = np.poly1d(z) + + # Calculate R-squared + r_squared = 1 - (sum((y - p(x))**2) / ((len(y) - 1) * np.var(y, ddof=1))) + + # Plot trendline + plt.plot(x, p(x), "r--", alpha=0.8, + label=f'Trendline: y={z[0]:.2f}x+{z[1]:.2f}\nR²: {r_squared:.4f}') + + plt.xlabel(param) + plt.ylabel('Score') + plt.title(f'Effect of {param} on Score') + + # Adjust x-axis labels + if len(values_list) > 10: + plt.xticks(range(0, len(values_list), len(values_list)//10), + [values_list[i] for i in range(0, len(values_list), len(values_list)//10)], + rotation=45, ha='right') + else: + plt.xticks(range(len(values_list)), values_list, rotation=45, ha='right') + + # Adjust legend + if len(values_list) > 10: + plt.legend(title="Legend", bbox_to_anchor=(1.05, 1), loc='upper left') + else: + plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') + + plt.tight_layout() + + # Save figure with error handling + try: + plt.savefig(f'{param}_vs_score.png', dpi=200, bbox_inches='tight') + except ValueError: + print(f"Warning: Failed to save image for {param}. skipping..") + + + plt.close() + + print("Plots have been saved as PNG files in the current directory.") + +if __name__ == "__main__": + root_dir = "/home/rednax/SSD2TB/Github_repos/diffusion_trainer/lora_models/PLANTOID" + + # Collect data + data = collect_data(root_dir) + + # Identify varying hyperparameters + varying_params = identify_varying_hyperparams(data) + + # Create plots + create_plots(data, varying_params) + + print("Plots have been saved as PNG files in the current directory.") \ No newline at end of file diff --git a/train_configs/test.json b/train_configs/test.json index bf62186..6003596 100644 --- a/train_configs/test.json +++ b/train_configs/test.json @@ -1,7 +1,7 @@ { "name": "xander_test", "sd_model_version": "sdxl", - "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip", + "lora_training_urls": "/home/rednax/Documents/datasets/01_people/mira", "concept_mode": "face", "seed": 1, "resolution": 512, @@ -10,9 +10,9 @@ "max_train_steps": 400, "augment_imgs_up_to_n": 40, "token_warmup_steps": 0, - "checkpointing_steps": 100, + "checkpointing_steps": 200, "ti_lr": 0.001, - "ti_weight_decay": 0.0005, + "ti_weight_decay": 0.00, "disable_ti": false, "n_tokens": 3, diff --git a/trainer/config.py b/trainer/config.py index bc274c7..269f68b 100644 --- a/trainer/config.py +++ b/trainer/config.py @@ -48,30 +48,30 @@ class TrainingConfig(BaseModel): train_img_size: List[int] = None train_aspect_ratio: float = None train_batch_size: int = 4 - num_train_epochs: int = 10000 max_train_steps: int = 360 + num_train_epochs: int = None checkpointing_steps: int = 10000 gradient_accumulation_steps: int = 1 is_lora: bool = True unet_optimizer_type: Literal["adamw", "prodigy", "AdamW8bit"] = "adamw" unet_lr_warmup_steps: int = None # slowly increase the learning rate of the adamw unet optimizer - unet_lr: float = 1.0e-3 + unet_lr: float = 0.0002 prodigy_d_coef: float = 1.0 unet_prodigy_growth_factor: float = 1.05 # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs) lora_weight_decay: float = 0.002 ti_lr: float = 1e-3 - ti_lr_warmup_steps: int = 20 # slowly ramp up the learning rate to build some momentum + ti_lr_warmup_steps: int = 10 # slowly ramp up the learning rate to build some momentum token_warmup_steps: int = 0 # warmup the token embeddings with a pure txt loss ti_weight_decay: float = 0.0 ti_optimizer: Literal["adamw", "prodigy"] = "adamw" freeze_ti_after_completion_f: float = 1.0 # freeze the TI after this fraction of the training is done + token_attention_loss_w: float = 4e-7 cond_reg_w: float = 0.0e-5 tok_cond_reg_w: float = 0.0e-5 - tok_cov_reg_w: float = 2000. # regularizes the token covariance matrix wrt pretrained "healthy" tokens - off_ratio_power: float = 0.02 # Pulls the std of the token distribution towards the target std + tok_cov_reg_w: float = 500. # regularizes the token covariance matrix wrt pretrained, normal tokens l1_penalty: float = 0.01 # Makes the unet lora matrix more sparse noise_offset: float = 0.02 # Noise offset training to improve very dark / very bright images @@ -92,17 +92,13 @@ class TrainingConfig(BaseModel): debug: bool = False allow_tf32: bool = True disable_ti: bool = False + skip_gpt_cleanup: bool = False weight_type: Literal["fp16", "bf16", "fp32"] = "bf16" n_tokens: int = 2 inserting_list_tokens: List[str] = ["",""] token_dict: dict = {"TOK": ""} device: str = "cuda:0" - crops_coords_top_left_h: int = 0 - crops_coords_top_left_w: int = 0 do_cache: bool = True - unet_learning_rate: float = 1.0 - lr_num_cycles: int = 1 - lr_power: float = 1.0 sample_imgs_lora_scale: float = None # Default lora scale for sampling the validation images dataloader_num_workers: int = 0 training_attributes: dict = {} @@ -129,13 +125,12 @@ class TrainingConfig(BaseModel): self.pretrained_model = {"path": self.ckpt_path, "url": None, "version": None} # add some metrics to the foldername: - lora_str = "dora" if self.use_dora else "lora" timestamp_short = datetime.now().strftime("%d_%H-%M-%S") if not self.name: - self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}" + self.name = "unnamed" - self.output_dir = self.output_dir + f"/{self.name}/" + f"{timestamp_short}-{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}_{self.max_train_steps}" + self.output_dir = self.output_dir + f"/{self.name}_" + f"{timestamp_short}-{self.concept_mode}_{self.resolution}_{self.caption_model}_{self.max_train_steps}" os.makedirs(self.output_dir, exist_ok=True) if self.seed is None: diff --git a/trainer/dataset.py b/trainer/dataset.py index 944fb0d..71e89b4 100644 --- a/trainer/dataset.py +++ b/trainer/dataset.py @@ -69,11 +69,9 @@ class PreprocessedDataset(Dataset): self.do_cache = True for idx in range(len(self.data)): - if len(self.data) < 25: - print(self.captions[idx]) vae_latent, mask = self._process(idx) self.vae_latents.append(vae_latent) - self.masks.append(mask) + self.masks.append(mask.detach()) print(f"\nCached latents, masks and captions for {len(self.vae_latents)} images.") del self.vae_encoder @@ -100,12 +98,9 @@ class PreprocessedDataset(Dataset): def get_aspect_ratio_bucketed_batch(self): assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ratio bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()" indices, resolution = self.bucket_manager.get_batch() - - print(f"Got bucket batch: {indices}, resolution: {resolution}") tok1, tok2, vae_latents, masks = [], [], [], [] for idx in indices: - if self.tokenizer_2 is None: t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution) else: @@ -141,28 +136,24 @@ class PreprocessedDataset(Dataset): image = PIL.Image.open(image_path).convert("RGB") if bucketing_resolution is None: image = prepare_image(image, w = self.size[0], h = self.size[1], pipe = self.pipe).to( - dtype=self.vae_encoder.dtype, device=self.vae_encoder.device + dtype=self.vae_encoder.dtype ) else: image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1], pipe = self.pipe).to( - dtype=self.vae_encoder.dtype, device=self.vae_encoder.device + dtype=self.vae_encoder.dtype ) - vae_latent = self.vae_encoder.encode(image).latent_dist + vae_latent = self.vae_encoder.encode(image.to(self.vae_encoder.device)).latent_dist dummy_vae_latent = vae_latent.sample() if self.mask_path is None: - mask = torch.ones_like( - dummy_vae_latent, dtype=self.vae_encoder.dtype, device=self.vae_encoder.device - ) + mask = torch.ones_like(dummy_vae_latent, dtype=self.vae_encoder.dtype) else: mask_path = self.mask_path[idx] mask_path = os.path.join(self.data_dir, mask_path) mask = PIL.Image.open(mask_path) - mask = prepare_mask(mask, self.size[0], self.size[1]).to( - dtype=self.vae_encoder.dtype, device=self.vae_encoder.device - ) + mask = prepare_mask(mask, self.size[0], self.size[1]).to(dtype=self.vae_encoder.dtype) mask_dtype = mask.dtype mask = mask.float() @@ -182,10 +173,10 @@ class PreprocessedDataset(Dataset): if self.do_cache: vae_latent = self.vae_latents[idx].sample() * self.vae_scaling_factor - return self.captions[idx], vae_latent.squeeze(), self.masks[idx] + return self.captions[idx], vae_latent.squeeze().detach(), self.masks[idx].detach() else: # This code pathway has not been tested in a long time and might be broken caption, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution) vae_latent = vae_latent.sample() * self.vae_scaling_factor - return caption, vae_latent.squeeze(), mask + return caption, vae_latent.squeeze().detach(), mask.detach() diff --git a/trainer/embedding_handler.py b/trainer/embedding_handler.py index ecd5ba7..ccec96f 100644 --- a/trainer/embedding_handler.py +++ b/trainer/embedding_handler.py @@ -260,11 +260,7 @@ class TokenEmbeddingsHandler: # original_size = (config.resolution, config.resolution) original_size = (1024, 1024) target_size = (config.resolution, config.resolution) - - crops_coords_top_left = ( - config.crops_coords_top_left_h, - config.crops_coords_top_left_w, - ) + crops_coords_top_left = (0,0) if pipe.text_encoder_2 is None: text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) @@ -397,7 +393,6 @@ class TokenEmbeddingsHandler: embedding_tensor.grad.data[:-config.n_tokens, : ] *= 0. optimizer_ti.step() - self.fix_embedding_std(config.off_ratio_power) optimizer_ti.zero_grad() if config.debug: @@ -430,41 +425,6 @@ class TokenEmbeddingsHandler: def device(self): return self.text_encoders[0].device - def fix_embedding_std(self, off_ratio_power=0.1): - if off_ratio_power == 0.0: - return - - idx = 0 - for tokenizer, text_encoder in zip(self.tokenizers, self.text_encoders): - if text_encoder is None: - idx += 1 - continue - - # Get the standard deviation target and current embeddings. - target_std = self.embeddings_settings[f"std_token_embedding_{idx}"] - embeddings, _ = self.get_trainable_embeddings() - new_embeddings = embeddings[f'txt_encoder_{idx}'] - assert len(new_embeddings.shape) == 2, "Embeddings should be 2D!" - - new_stds = new_embeddings.std(dim=1) - #off_ratios = target_std.float() / new_stds.float() - off_ratios = target_std / new_stds - - # Check if off_ratios are within an acceptable range. - if (off_ratios.min() < 0.9) or (off_ratios.max() > 1.1): - # Convert the pytorch tensor into a list of python floats: - off_ratio_float_list = np.round(off_ratios.detach().float().cpu().numpy().tolist(), 3) - print(f"WARNING: std-off ratio-{idx} (target-std / embedding-std) token-ratios = {off_ratio_float_list}, prob not ideal...") - - # Adjust embeddings using the computed ratios. - index_no_updates = self.embeddings_settings[f"index_no_updates_{idx}"] - index_updates = ~index_no_updates - multiplier_values = off_ratios**off_ratio_power - multiplier_values = multiplier_values.unsqueeze(1).expand_as(new_embeddings) - text_encoder.text_model.embeddings.token_embedding.weight.data[index_updates] *= multiplier_values - - idx += 1 - def _load_embeddings(self, loaded_embeddings, tokenizer, text_encoder): # Assuming new tokens are of the format self.inserting_toks = [f"" for i in range(loaded_embeddings.shape[0])] diff --git a/trainer/inference.py b/trainer/inference.py index 58d66ff..7b0ffad 100644 --- a/trainer/inference.py +++ b/trainer/inference.py @@ -155,11 +155,7 @@ def get_conditioning_signals(config, pipe, captions): # original_size = (config.resolution, config.resolution) original_size = (1024, 1024) target_size = (config.resolution, config.resolution) - - crops_coords_top_left = ( - config.crops_coords_top_left_h, - config.crops_coords_top_left_w, - ) + crops_coords_top_left = (0,0) if pipe.text_encoder_2 is None: text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) diff --git a/trainer/loss.py b/trainer/loss.py index 5fc9ff2..728f61c 100644 --- a/trainer/loss.py +++ b/trainer/loss.py @@ -6,72 +6,72 @@ from torch.utils._foreach_utils import _group_tensors_by_device_and_dtype, _has_ from trainer.inference import get_conditioning_signals import torch.nn.functional as F -def compute_token_attention_loss(pipe, embedding_handler, - captions, masks, - daam_loss, verbose = 0 - ): + +def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_loss, verbose=0): """ - distribution shift loss + Custom loss function to regularize the attention maps of the token embeddings. """ + masks = masks[:, 0].float() + img_ratio = masks.shape[-1] / masks.shape[-2] - ti_heatmaps, ti_masks, att_L2_losses = [], [], [] + att_L2_losses = [] + ti_heatmaps = [] + ti_masks = [] att_reg_threshold = -5.0 - input_dtype = masks.dtype - masks = masks[:,0,:,:].float() + # attention_maps.shape = [n_layers, batch_size, w, h, 77] + attention_maps = daam_loss.process_and_stack_attention_scores(img_ratio) + n_layers, batch_size, w, h, n_tokens = attention_maps.shape - img_ratio = masks.shape[-1] / masks.shape[-2] + # reshape masks to match attention maps: + # masks.shape = [batch_size, w2, h2] + masks = F.interpolate(masks.unsqueeze(1), size=(attention_maps.shape[-3], attention_maps.shape[-2])).squeeze(1) + masks = masks.unsqueeze(0).unsqueeze(-1) + masks = masks.repeat(n_layers, 1, 1, 1, n_tokens) - for batch_index in range(len(captions)): - mask = masks[batch_index] + for batch_index, caption in enumerate(captions): + token_indices = pipe.tokenizer.encode(caption) - #token_strings = [ - # pipe.tokenizer.decode(x) - # for x in pipe.tokenizer.encode(captions[batch_index]) - # ] - token_indices_in_prompt = pipe.tokenizer.encode(captions[batch_index]) - - mean_att_per_token = daam_loss.get_mean_attention_per_token(token_indices_in_prompt, batch_index) + # Penalize the mean attention score of each token: + mean_att_per_token = attention_maps[:,batch_index, :, :, 1:len(token_indices)-1].mean(dim=[0,1,2]) att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean() att_L2_losses.append(att_L2_loss) - - trained_token_indices = embedding_handler.train_ids - # Find the index of the trained tokens in token_indices_in_prompt: - ti_token_indices = [ - token_indices_in_prompt.index(token_index) - for token_index in trained_token_indices - ] + ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids] + batch_ti_heatmaps, batch_ti_masks = [], [] + # Extract the attention heatmaps corresponding to the trainable token embeddings: for text_token_index in ti_token_indices: - ti_heatmap = daam_loss.get_the_daam_heatmap(text_token_index = text_token_index, img_ratio=img_ratio)[batch_index] - resized_mask = F.interpolate(input = mask.unsqueeze(0).unsqueeze(0), size = (ti_heatmap.shape[-2], ti_heatmap.shape[-1])).squeeze(0).squeeze(0) + ti_heatmap = attention_maps[:,batch_index, :, :, text_token_index].mean(dim=0) + ti_mask = masks[:,batch_index, :, :, text_token_index].mean(dim=0) + batch_ti_heatmaps.append(ti_heatmap.float()) + batch_ti_masks.append(ti_mask) - ti_heatmaps.append(ti_heatmap.float()) - ti_masks.append(resized_mask.squeeze(0)) + ti_heatmaps.append(torch.stack(batch_ti_heatmaps)) + ti_masks.append(torch.stack(batch_ti_masks)) - # ti_heatmaps shape = [n_tokens x batch_size, w, h] ti_heatmaps = torch.stack(ti_heatmaps) ti_masks = torch.stack(ti_masks) - # Penalize large, positive mean token attentions: - reg_loss_0 = 20*torch.stack(att_L2_losses).mean() + #ti_heatmaps.shape = [batch_size, n_tokens, w, h] + token_means = ti_heatmaps.mean(dim=[2,3]) + token_attention_scores = token_means.var(dim=1) - # Where the segmentation mask is one, we want to avoid very large attention scores: - threshold = 0.0 - reg_loss_1 = 0.25*(torch.relu(ti_heatmaps*ti_masks - threshold)**2).mean() - - # Where the segmentation mask is zero, we want to avoid somewhat large attention scores: - threshold = -10.0 - reg_loss_2 = 0.1 * (torch.relu(ti_heatmaps*(1.0-ti_masks) - threshold)**2).mean() + # Avoid large attention scores in general: + reg_loss_0 = 5.0 * torch.stack(att_L2_losses).mean() + # Avoid large attention scores for ti tokens, inside the masked region: + reg_loss_1 = 1.0 * (torch.relu(ti_heatmaps * ti_masks)**2).mean() + # Avoid large attention scores for ti tokens, outside of the masked region: + reg_loss_2 = 0.5 * (torch.relu(ti_heatmaps * (1 - ti_masks) + 10)**2).mean() + # Make the Ti tokens have equal avg attention scores (equal distribution of concept information): + reg_loss_3 = 10.0 * token_attention_scores.mean() if verbose: print(f"reg_loss_0: {reg_loss_0.item():.4f}") print(f"reg_loss_1: {reg_loss_1.item():.4f}") print(f"reg_loss_2: {reg_loss_2.item():.4f}") + print(f"reg_loss_3: {reg_loss_3.item():.4f}") - reg_loss = reg_loss_0 + reg_loss_1 + reg_loss_2 - - return reg_loss.to(input_dtype) + return (reg_loss_0 + reg_loss_1 + reg_loss_2 + reg_loss_3).to(masks.dtype) def compute_snr(noise_scheduler, timesteps): diff --git a/trainer/models.py b/trainer/models.py index 7bc90b8..9f165ba 100644 --- a/trainer/models.py +++ b/trainer/models.py @@ -4,24 +4,30 @@ import subprocess import torch from diffusers import AutoencoderKL, DDPMScheduler, EulerDiscreteScheduler, UNet2DConditionModel, StableDiffusionPipeline, StableDiffusionXLPipeline -def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae_float32 = False): +def load_models(pretrained_model, device, weight_dtype = torch.float16): # check if the model is already downloaded: if not os.path.exists(pretrained_model['path']): download_weights(pretrained_model['url'], pretrained_model['path']) + tokenizer_two, text_encoder_two = None, None print(f"Loading model weights from {os.path.abspath(pretrained_model['path'])} with dtype: {weight_dtype}...") try: + print("Loading as SDXL model...") pipe = StableDiffusionXLPipeline.from_single_file( pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True) sd_model_version = "sdxl" + tokenizer_two = pipe.tokenizer_2 + text_encoder_two = pipe.text_encoder_2 + text_encoder_two.requires_grad_(False) + text_encoder_two.to(device, dtype=weight_dtype) except: + print("Loading as SD15 model...") pipe = StableDiffusionPipeline.from_single_file( pretrained_model['path'], torch_dtype=weight_dtype, use_safetensors=True) sd_model_version = "sd15" print(f"Loaded {sd_model_version} model!") - pipe = pipe.to(device, dtype=weight_dtype) noise_scheduler = DDPMScheduler.from_config(pipe.scheduler.config) @@ -31,24 +37,11 @@ def load_models(pretrained_model, device, weight_dtype = torch.float16, keep_vae text_encoder_one = pipe.text_encoder vae.requires_grad_(False) - if keep_vae_float32: - vae.to(device, dtype=torch.float32) - else: - vae.to(device, dtype=weight_dtype) - if weight_dtype != torch.float32: - print(f"Warning: VAE will be loaded as {weight_dtype}, this is fine for inference but may not be ideal for training..?") - + vae.to(device, dtype=weight_dtype) unet.to(device, dtype=weight_dtype) text_encoder_one.requires_grad_(False) text_encoder_one.to(device, dtype=weight_dtype) - - tokenizer_two = text_encoder_two = None - if sd_model_version == "sdxl": - tokenizer_two = pipe.tokenizer_2 - text_encoder_two = pipe.text_encoder_2 - text_encoder_two.requires_grad_(False) - text_encoder_two.to(device, dtype=weight_dtype) - + return ( pipe, tokenizer_one, diff --git a/trainer/preprocess.py b/trainer/preprocess.py index e3357d4..fb6b489 100755 --- a/trainer/preprocess.py +++ b/trainer/preprocess.py @@ -331,12 +331,12 @@ def extract_gpt_concept_description(gpt_completion, concept_mode): return concept_name -def post_process_captions(captions, text, concept_mode, job_seed): +def post_process_captions(captions, text, concept_mode, job_seed, skip_gpt_cleanup=False): text = text.strip() gpt_cleanup_worked = False gpt_concept_description = None - if len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client: + if (len(captions) >= MIN_GPT_PROMPTS and len(captions) <= MAX_GPT_PROMPTS and not text and client) and not skip_gpt_cleanup: retry_count = 0 while retry_count < 5: try: @@ -841,14 +841,13 @@ def load_and_save_masks_and_captions( trigger_text = "" gpt_concept_description = None if not config.disable_ti: - captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed) + captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed, skip_gpt_cleanup=config.skip_gpt_cleanup) if (gpt_concept_description is not None) and ((mask_target_prompts is None) or (mask_target_prompts == "")): print(f"Using GPT concept name as CLIP-segmentation prompt: {gpt_concept_description}") mask_target_prompts = gpt_concept_description if mask_target_prompts is None or config.concept_mode == "style": - print("Disabling CLIP-segmentation") mask_target_prompts = "" temp = 999 else: @@ -923,10 +922,6 @@ def load_and_save_masks_and_captions( else: captions = ["TOK, " + caption if "TOK" not in caption else caption for caption in captions] - print("Final captions:") - for caption in captions: - print(caption) - # iterate through the images, masks, and captions and add a row to the dataframe for each print("Saving final training dataset...") for idx, (image, mask, caption) in enumerate(zip(images, seg_masks, captions)): diff --git a/trainer/ti_cross_attn_loss.py b/trainer/ti_cross_attn_loss.py index 01c3d75..8ea00d9 100644 --- a/trainer/ti_cross_attn_loss.py +++ b/trainer/ti_cross_attn_loss.py @@ -234,6 +234,36 @@ class DAAMLoss: x.name for x in attention_processors ] + def process_and_stack_attention_scores(self, img_ratio: float): + reshaped_tensors = [] + min_heatmap_pixels = np.inf + + # Process each attention score + for processor in self.attention_processors: + score = processor.cross_attention_scores + bs, seq_len, channels = score.shape + + # Calculate width and height based on img_ratio + width = round(math.sqrt(seq_len * img_ratio)) + height = round(width / img_ratio) + reshaped_score = rearrange(score, 'b (h w) c -> b h w c', h=height, w=width) + reshaped_tensors.append(reshaped_score) + min_heatmap_pixels = min(min_heatmap_pixels, height * width) + + final_height = round(math.sqrt(min_heatmap_pixels)) + final_width = round(final_height / img_ratio) + + # Interpolate and standardize all tensors to the same size + for i, heatmap in enumerate(reshaped_tensors): + if heatmap.shape[1] * heatmap.shape[2] != min_heatmap_pixels: + # Interpolating to match the smallest tensor size uniformly + heatmap = F.interpolate(heatmap.permute(0, 3, 1, 2), size=(final_height, final_width), mode='bilinear').permute(0, 2, 3, 1) + reshaped_tensors[i] = heatmap + + # Stack all tensors along the first dimension + stacked_tensor = torch.stack(reshaped_tensors, dim=0) + return stacked_tensor + def get_all_cross_attention_scores(self): cross_attention_scores = {} @@ -243,20 +273,6 @@ class DAAMLoss: ] = p.cross_attention_scores return cross_attention_scores - - def get_mean_attention_per_token(self, token_indices_in_prompt, batch_index): - cross_attention_scores = self.get_all_cross_attention_scores() - - means = [] - - for layer_name in cross_attention_scores.keys(): - attention = cross_attention_scores[layer_name][batch_index][:, 1:len(token_indices_in_prompt)-1] - attention_mean_per_token = attention.mean(dim=0) - means.append(attention_mean_per_token) - - mean_attention_per_token = torch.stack(means).mean(dim=0) - - return mean_attention_per_token def get_image_heatmap(self, text_token_index: int, layer_name: str, img_ratio: float) -> TensorType["batch", "height", "width"]: cross_attention_scores = self.get_all_cross_attention_scores() @@ -313,10 +329,8 @@ def get_module_by_name(module: nn.Module, name: str): names = name.split(sep=".") return reduce(getattr, names, module) -def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionXLPipeline, DAAMLoss]: - - assert isinstance(pipeline, StableDiffusionXLPipeline) +def init_daam_loss(pipeline): ## find out where the attention processor thingies are module_names = find_attnprocessor2_0( unet = pipeline.unet @@ -325,8 +339,6 @@ def init_daam_loss(pipeline: StableDiffusionXLPipeline)-> tuple[StableDiffusionX all_daam_attention_processors = [] # override the attention processor thingies for name in module_names: - # print(f"Replacing: {name}") - # Get parent module and attribute name parent_name = ".".join(name.split(".")[:-1]) attr_name = name.split(".")[-1]