From e9555cd5842fb8a23e52d4fe76bc34cbbb57184b Mon Sep 17 00:00:00 2001 From: xander Date: Mon, 12 Aug 2024 14:13:10 +0200 Subject: [PATCH] updates --- debug/generate_test_imgs.py | 28 ------ main.py | 5 + train_configs/training_args_face_sdxl.json | 6 +- train_configs/training_args_style_sdxl.json | 4 +- trainer/config.py | 1 + trainer/loss.py | 8 +- trainer/ti_cross_attn_loss.py | 101 ++++++++++---------- 7 files changed, 71 insertions(+), 82 deletions(-) delete mode 100755 debug/generate_test_imgs.py diff --git a/debug/generate_test_imgs.py b/debug/generate_test_imgs.py deleted file mode 100755 index d756e65..0000000 --- a/debug/generate_test_imgs.py +++ /dev/null @@ -1,28 +0,0 @@ -import numpy as np -from PIL import Image -import os - -def generate_random_color_image(width, height): - """Generate an image of random color.""" - color = np.random.randint(0, 256, (3,), dtype=np.uint8) - image = np.full((height, width, 3), color, dtype=np.uint8) - return Image.fromarray(image) - -def save_images(num_images, width, height, directory): - """Save a specified number of random color images.""" - for i in range(num_images): - image = generate_random_color_image(width, height) - image.save(f"{directory}/random_color_image_{i+1}.png") - -# Parameters -num_images = 40 -width = 1024 -height = 1024 -directory = "random_images" - -os.makedirs(directory, exist_ok=True) - -# Generate and save images -save_images(num_images, width, height, directory) - -print(f'Saved {num_images} random color images to {os.path.abspath(directory)}') \ No newline at end of file diff --git a/main.py b/main.py index 5451d8a..88cac97 100755 --- a/main.py +++ b/main.py @@ -291,6 +291,11 @@ def train(config: TrainingConfig): mask = mask.to(config.device) captions = list(captions) + if config.caption_dropout > 0.0: + for i in range(len(captions)): + if np.random.rand() < config.caption_dropout: + captions[i] = "" + prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals( config, pipe, captions ) diff --git a/train_configs/training_args_face_sdxl.json b/train_configs/training_args_face_sdxl.json index 1d1a3f1..66e6867 100644 --- a/train_configs/training_args_face_sdxl.json +++ b/train_configs/training_args_face_sdxl.json @@ -3,7 +3,7 @@ "sd_model_version": "sdxl", "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip", "concept_mode": "face", - "sample_imgs_lora_scale": 0.75, + "sample_imgs_lora_scale": 0.8, "seed": 0, "resolution": 512, "train_batch_size": 4, @@ -11,8 +11,10 @@ "max_train_steps": 250, "token_warmup_steps": 0, "checkpointing_steps": 175, + + "disable_ti": false, - "ti_lr": 0.001, + "ti_lr": 0.0005, "unet_lr": 0.001, "lora_rank": 16, diff --git a/train_configs/training_args_style_sdxl.json b/train_configs/training_args_style_sdxl.json index 81a6626..e05e739 100644 --- a/train_configs/training_args_style_sdxl.json +++ b/train_configs/training_args_style_sdxl.json @@ -1,7 +1,7 @@ { - "name": "twisting_realities", + "name": "clipx", "sd_model_version": "sdxl", - "lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/twisting_realities.zip", + "lora_training_urls": "https://edenartlab-lfs.s3.amazonaws.com/datasets/clipx.zip", "concept_mode": "style", "sample_imgs_lora_scale": 0.8, "seed": 0, diff --git a/trainer/config.py b/trainer/config.py index 38e1b75..ba031ca 100644 --- a/trainer/config.py +++ b/trainer/config.py @@ -38,6 +38,7 @@ class TrainingConfig(BaseModel): concept_mode: Literal["face", "style", "object"] caption_prefix: str = "" # hardcoding this will inject TOK manually and skip the chatgpt token injection step, not recommended unless you know what you're doing caption_model: Literal["gpt4-v", "blip", "florence"] = "blip" + caption_dropout: float = 0.0 # dropout rate for captions: occasionally use empty prompt sd_model_version: Literal["sdxl", "sd15", None] = None ckpt_path: str = None # optional hardcoded checkpoint path pretrained_model: dict = None diff --git a/trainer/loss.py b/trainer/loss.py index 728f61c..e5c621f 100644 --- a/trainer/loss.py +++ b/trainer/loss.py @@ -37,7 +37,10 @@ def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_ att_L2_loss = (torch.relu(mean_att_per_token - att_reg_threshold)**2).mean() att_L2_losses.append(att_L2_loss) - ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids] + try: + ti_token_indices = [token_indices.index(token_id) for token_id in embedding_handler.train_ids] + except: + continue batch_ti_heatmaps, batch_ti_masks = [], [] # Extract the attention heatmaps corresponding to the trainable token embeddings: for text_token_index in ti_token_indices: @@ -49,6 +52,9 @@ def compute_token_attention_loss(pipe, embedding_handler, captions, masks, daam_ ti_heatmaps.append(torch.stack(batch_ti_heatmaps)) ti_masks.append(torch.stack(batch_ti_masks)) + if len(ti_heatmaps) == 0: + return torch.tensor(0.0).to(masks.dtype) + ti_heatmaps = torch.stack(ti_heatmaps) ti_masks = torch.stack(ti_masks) diff --git a/trainer/ti_cross_attn_loss.py b/trainer/ti_cross_attn_loss.py index bd49a00..f1b5b08 100644 --- a/trainer/ti_cross_attn_loss.py +++ b/trainer/ti_cross_attn_loss.py @@ -17,67 +17,70 @@ import numpy as np from mpl_toolkits.axes_grid1 import make_axes_locatable def plot_token_attention_loss(folder, pipe, daam_loss, captions, timesteps, token_attention_loss, global_step, img_ratio): - batch_index = 0 - timestep = timesteps[batch_index].item() - if timestep > 700: - batch_index += 1 + try: + batch_index = 0 timestep = timesteps[batch_index].item() + if timestep > 700: + batch_index += 1 + timestep = timesteps[batch_index].item() - folder = os.path.join(folder, "attention_heatmaps") - os.makedirs(folder, exist_ok=True) + folder = os.path.join(folder, "attention_heatmaps") + os.makedirs(folder, exist_ok=True) - token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])] - plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens + token_strings = [pipe.tokenizer.decode(x) for x in pipe.tokenizer.encode(captions[batch_index])] + plot_token_indices = range(1, len(token_strings) - 1) # Exclude start and end tokens - # Calculate global min and max for consistent colormap - all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float() - for i in plot_token_indices] - vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps]) + # Calculate global min and max for consistent colormap + all_heatmaps = [daam_loss.get_the_daam_heatmap(text_token_index=i, img_ratio=img_ratio)[batch_index].cpu().detach().float() + for i in plot_token_indices] + vmin, vmax = np.min([h.min() for h in all_heatmaps]), np.max([h.max() for h in all_heatmaps]) - # Heatmap plots - fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices), - figsize=(3 * len(plot_token_indices), 10)) - title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}" - fig.suptitle(title_str, fontsize=16) - - for idx, text_token_index in enumerate(plot_token_indices): - heatmap = all_heatmaps[idx] - im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax) - axes[idx].set_title(f"{token_strings[text_token_index]}") - axes[idx].axis("off") + # Heatmap plots + fig, axes = plt.subplots(nrows=1, ncols=len(plot_token_indices), + figsize=(3 * len(plot_token_indices), 10)) + title_str = f"Token Attention Heatmaps (Step: {global_step})\nDenoise Timestep: {timesteps[batch_index].item()}" + fig.suptitle(title_str, fontsize=16) - # Add colorbar - divider = make_axes_locatable(axes[idx]) - cax = divider.append_axes("bottom", size="5%", pad=0.05) - plt.colorbar(im, cax=cax, orientation="horizontal") + for idx, text_token_index in enumerate(plot_token_indices): + heatmap = all_heatmaps[idx] + im = axes[idx].imshow(heatmap, cmap='viridis', vmin=vmin, vmax=vmax) + axes[idx].set_title(f"{token_strings[text_token_index]}") + axes[idx].axis("off") + + # Add colorbar + divider = make_axes_locatable(axes[idx]) + cax = divider.append_axes("bottom", size="5%", pad=0.05) + plt.colorbar(im, cax=cax, orientation="horizontal") - plt.tight_layout() - fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight') - plt.close(fig) + plt.tight_layout() + fig.savefig(os.path.join(folder, f"heatmaps_{global_step}.jpg"), dpi=300, bbox_inches='tight') + plt.close(fig) - # Histogram plot - fig, ax = plt.subplots(figsize=(12, 6)) - ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16) - ax.set_xlabel("Attention Value", fontsize=12) - ax.set_ylabel("Frequency", fontsize=12) + # Histogram plot + fig, ax = plt.subplots(figsize=(12, 6)) + ax.set_title(f"Token Attention Distribution (Step: {global_step})", fontsize=16) + ax.set_xlabel("Attention Value", fontsize=12) + ax.set_ylabel("Frequency", fontsize=12) - for idx, text_token_index in enumerate(plot_token_indices): - heatmap = all_heatmaps[idx] - ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index], - alpha=0.5, density=True) + for idx, text_token_index in enumerate(plot_token_indices): + heatmap = all_heatmaps[idx] + ax.hist(heatmap.reshape(-1), bins=30, label=token_strings[text_token_index], + alpha=0.5, density=True) - ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10) - ax.grid(alpha=0.3) - plt.tight_layout() + ax.legend(bbox_to_anchor=(1.05, 1), loc='upper left', fontsize=10) + ax.grid(alpha=0.3) + plt.tight_layout() - # Add token attention loss as text - plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}", - transform=ax.transAxes, ha='right', va='top', - bbox=dict(facecolor='white', edgecolor='black', alpha=0.8)) + # Add token attention loss as text + plt.text(0.95, 0.95, f"Token Attention Loss: {token_attention_loss.item():.4f}", + transform=ax.transAxes, ha='right', va='top', + bbox=dict(facecolor='white', edgecolor='black', alpha=0.8)) - plt.ylim(0, 0.6) - fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight') - plt.close(fig) + plt.ylim(0, 0.6) + fig.savefig(os.path.join(folder, f"histogram_{global_step}.jpg"), dpi=300, bbox_inches='tight') + plt.close(fig) + except: + print("Failed to plot token attention loss")