diff --git a/.gitignore b/.gitignore index 6cf35b9..a3661b5 100644 --- a/.gitignore +++ b/.gitignore @@ -17,6 +17,7 @@ gridsearch* aesthetic_score_best_model.pth # experiment folders: +scripts/plots conditioning_spaces/ training_args_x_*.json xander_configs/ diff --git a/faces.sh b/faces.sh deleted file mode 100644 index abbbd54..0000000 --- a/faces.sh +++ /dev/null @@ -1,87 +0,0 @@ -#!/bin/bash - -python main.py scripts/gridsearch_configs/faces/faces_000.json -python main.py scripts/gridsearch_configs/faces/faces_001.json -python main.py scripts/gridsearch_configs/faces/faces_002.json -python main.py scripts/gridsearch_configs/faces/faces_003.json -python main.py scripts/gridsearch_configs/faces/faces_004.json -python main.py scripts/gridsearch_configs/faces/faces_005.json -python main.py scripts/gridsearch_configs/faces/faces_006.json -python main.py scripts/gridsearch_configs/faces/faces_007.json -python main.py scripts/gridsearch_configs/faces/faces_008.json -python main.py scripts/gridsearch_configs/faces/faces_009.json -python main.py scripts/gridsearch_configs/faces/faces_010.json -python main.py scripts/gridsearch_configs/faces/faces_011.json -python main.py scripts/gridsearch_configs/faces/faces_012.json -python main.py scripts/gridsearch_configs/faces/faces_013.json -python main.py scripts/gridsearch_configs/faces/faces_014.json -python main.py scripts/gridsearch_configs/faces/faces_015.json -python main.py scripts/gridsearch_configs/faces/faces_016.json -python main.py scripts/gridsearch_configs/faces/faces_017.json -python main.py scripts/gridsearch_configs/faces/faces_018.json -python main.py scripts/gridsearch_configs/faces/faces_019.json -python main.py scripts/gridsearch_configs/faces/faces_020.json -python main.py scripts/gridsearch_configs/faces/faces_021.json -python main.py scripts/gridsearch_configs/faces/faces_022.json -python main.py scripts/gridsearch_configs/faces/faces_023.json -python main.py scripts/gridsearch_configs/faces/faces_024.json -python main.py scripts/gridsearch_configs/faces/faces_025.json -python main.py scripts/gridsearch_configs/faces/faces_026.json -python main.py scripts/gridsearch_configs/faces/faces_027.json -python main.py scripts/gridsearch_configs/faces/faces_028.json -python main.py scripts/gridsearch_configs/faces/faces_029.json -python main.py scripts/gridsearch_configs/faces/faces_030.json -python main.py scripts/gridsearch_configs/faces/faces_031.json -python main.py scripts/gridsearch_configs/faces/faces_032.json -python main.py scripts/gridsearch_configs/faces/faces_033.json -python main.py scripts/gridsearch_configs/faces/faces_034.json -python main.py scripts/gridsearch_configs/faces/faces_035.json -python main.py scripts/gridsearch_configs/faces/faces_036.json -python main.py scripts/gridsearch_configs/faces/faces_037.json -python main.py scripts/gridsearch_configs/faces/faces_038.json -python main.py scripts/gridsearch_configs/faces/faces_039.json -python main.py scripts/gridsearch_configs/faces/faces_040.json -python main.py scripts/gridsearch_configs/faces/faces_041.json -python main.py scripts/gridsearch_configs/faces/faces_042.json -python main.py scripts/gridsearch_configs/faces/faces_043.json -python main.py scripts/gridsearch_configs/faces/faces_044.json -python main.py scripts/gridsearch_configs/faces/faces_045.json -python main.py scripts/gridsearch_configs/faces/faces_046.json -python main.py scripts/gridsearch_configs/faces/faces_047.json -python main.py scripts/gridsearch_configs/faces/faces_048.json -python main.py scripts/gridsearch_configs/faces/faces_049.json -python main.py scripts/gridsearch_configs/faces/faces_050.json -python main.py scripts/gridsearch_configs/faces/faces_051.json -python main.py scripts/gridsearch_configs/faces/faces_052.json -python main.py scripts/gridsearch_configs/faces/faces_053.json -python main.py scripts/gridsearch_configs/faces/faces_054.json -python main.py scripts/gridsearch_configs/faces/faces_055.json -python main.py scripts/gridsearch_configs/faces/faces_056.json -python main.py scripts/gridsearch_configs/faces/faces_057.json -python main.py scripts/gridsearch_configs/faces/faces_058.json -python main.py scripts/gridsearch_configs/faces/faces_059.json -python main.py scripts/gridsearch_configs/faces/faces_060.json -python main.py scripts/gridsearch_configs/faces/faces_061.json -python main.py scripts/gridsearch_configs/faces/faces_062.json -python main.py scripts/gridsearch_configs/faces/faces_063.json -python main.py scripts/gridsearch_configs/faces/faces_064.json -python main.py scripts/gridsearch_configs/faces/faces_065.json -python main.py scripts/gridsearch_configs/faces/faces_066.json -python main.py scripts/gridsearch_configs/faces/faces_067.json -python main.py scripts/gridsearch_configs/faces/faces_068.json -python main.py scripts/gridsearch_configs/faces/faces_069.json -python main.py scripts/gridsearch_configs/faces/faces_070.json -python main.py scripts/gridsearch_configs/faces/faces_071.json -python main.py scripts/gridsearch_configs/faces/faces_072.json -python main.py scripts/gridsearch_configs/faces/faces_073.json -python main.py scripts/gridsearch_configs/faces/faces_074.json -python main.py scripts/gridsearch_configs/faces/faces_075.json -python main.py scripts/gridsearch_configs/faces/faces_076.json -python main.py scripts/gridsearch_configs/faces/faces_077.json -python main.py scripts/gridsearch_configs/faces/faces_078.json -python main.py scripts/gridsearch_configs/faces/faces_079.json -python main.py scripts/gridsearch_configs/faces/faces_080.json -python main.py scripts/gridsearch_configs/faces/faces_081.json -python main.py scripts/gridsearch_configs/faces/faces_082.json -python main.py scripts/gridsearch_configs/faces/faces_083.json -python main.py scripts/gridsearch_configs/faces/faces_084.json diff --git a/scripts/create_hyperparam_sweep.py b/scripts/create_hyperparam_sweep.py index d27bec7..2e10089 100644 --- a/scripts/create_hyperparam_sweep.py +++ b/scripts/create_hyperparam_sweep.py @@ -66,15 +66,15 @@ hyperparameters = { "checkpointing_steps": [360], "gradient_accumulation_steps": [1], - "n_tokens": [2,3,4], - "ti_lr": [0.001, 0.003], + "n_tokens": [3], + "ti_lr": [0.001], "ti_weight_decay": [0.000], "l1_penalty": [0.0], - "token_warmup_steps": [0], + "token_warmup_steps": [0,60], "tok_cov_reg_w": [500], - "token_attention_loss_w": [0, 2e-7, 10e-7], + "token_attention_loss_w": [2e-7], - "unet_lr": [0.001, 0.002], + "unet_lr": [0.001, 0.0003], "lora_alpha_multiplier": [1.0], "prodigy_d_coef": [1.0], "lora_weight_decay": [0.001], @@ -88,7 +88,7 @@ hyperparameters = { "text_encoder_lora_lr": [0.0e-4], "snr_gamma": [5.0], - "caption_model": ["blip", "florence"], + "caption_model": ["florence"], "augment_imgs_up_to_n": [40], "verbose": ['true'], "debug": ['true'] diff --git a/scripts/eval_hyperparam_sweep.py b/scripts/eval_hyperparam_sweep.py index 2126dad..5f3bc4b 100644 --- a/scripts/eval_hyperparam_sweep.py +++ b/scripts/eval_hyperparam_sweep.py @@ -80,7 +80,7 @@ def identify_varying_hyperparams(data, skip_params=['output_dir', 'start_time', return varying_params -def create_plots(data, varying_params, outdir): +def create_plots(data, varying_params, outdir, top = 0.15): os.makedirs(outdir, exist_ok=True) for param, values in varying_params.items(): @@ -90,12 +90,17 @@ def create_plots(data, varying_params, outdir): plt.figure(figsize=(12, 8)) param_data = defaultdict(list) + all_scores = [] for args, score in data: if param in args: value = args[param] value_str = str(value) param_data[value_str].append(score) + all_scores.append(score) + + # Calculate global top threshold + global_top_percent = np.percentile(all_scores, 100*(1-top)) # Sort the values try: @@ -109,16 +114,15 @@ def create_plots(data, varying_params, outdir): for i, value_str in enumerate(values_list): scores = param_data[value_str] - # Apply jitter first + # Apply jitter jittered_x = np.random.normal(i, 0.1, size=len(scores)) jittered_y = np.array(scores) + np.random.normal(0, 0.01 * max(scores), size=len(scores)) - # Calculate top 25% based on ORIGINAL scores (before jitter) - top_25_percent = np.percentile(scores, 75) - top_25_mask = np.array(scores) >= top_25_percent + # Use global top 25% threshold + top_mask = np.array(scores) >= global_top_percent # Emphasize top 25% scores using JITTERED coordinates for plotting - sns.scatterplot(x=jittered_x[top_25_mask], y=jittered_y[top_25_mask], alpha=0.6, color='black', marker='X', s=40, linewidth=1) + sns.scatterplot(x=jittered_x[top_mask], y=jittered_y[top_mask], alpha=0.6, color='black', marker='X', s=40, linewidth=1) # Plot all scores using JITTERED coordinates sns.scatterplot(x=jittered_x, y=jittered_y, alpha=0.6, label=value_str) @@ -135,11 +139,10 @@ def create_plots(data, varying_params, outdir): plt.plot(range(len(values_list)), p(range(len(values_list))), "r--", alpha=0.8, label=f'All data: y={z[0]:.2f}x+{z[1]:.2f}\nR²: {r2_score(y, p(x)):.4f}') - # Calculate trendline for top 25% scoring datapoints - top_25_percent = np.percentile(y, 75) - top_25_mask = y >= top_25_percent - x_top = x[top_25_mask] - y_top = y[top_25_mask] + # Calculate trendline for global top 25% scoring datapoints + top_mask = y >= global_top_percent + x_top = x[top_mask] + y_top = y[top_mask] z_top = np.polyfit(x_top, y_top, 1) p_top = np.poly1d(z_top) @@ -169,7 +172,7 @@ def create_plots(data, varying_params, outdir): print(f"Plots have been saved as PNG files in {outdir}") if __name__ == "__main__": - root_dir = "/home/rednax/SSD2TB/Github_repos/diffusion_trainer/lora_models/OBJECTS" + root_dir = "/home/rednax/SSD2TB/Github_repos/diffusion_trainer/lora_models/faces" outdir = os.path.join('.', os.path.basename(root_dir)) # Collect data diff --git a/train_configs/training_args_face_sdxl.json b/train_configs/training_args_face_sdxl.json index 7552ef7..35653e4 100644 --- a/train_configs/training_args_face_sdxl.json +++ b/train_configs/training_args_face_sdxl.json @@ -1,27 +1,17 @@ { - "name": "xander_sdxl", + "name": "gene_sdxl", "sd_model_version": "sdxl", - "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip", + "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip", "concept_mode": "face", - "seed": 1, + "seed": 0, "resolution": 512, "train_batch_size": 4, "n_sample_imgs": 6, - "max_train_steps": 400, - "token_warmup_steps": 0, - "checkpointing_steps": 200, + "max_train_steps": 360, + "checkpointing_steps": 180, "ti_lr": 0.001, - "ti_weight_decay": 0.0005, - - "disable_ti": false, - "text_encoder_lora_optimizer": null, - "text_encoder_lora_lr": 1.0e-4, - "text_encoder_lora_weight_decay": 1e-5, - "text_encoder_lora_rank": 12, - "unet_lr": 0.001, "lora_rank": 16, - "use_dora": false, - "caption_model": "blip", + "debug": true } \ No newline at end of file diff --git a/train_configs/training_args_style_sdxl.json b/train_configs/training_args_style_sdxl.json index 2dc9884..6d78f01 100644 --- a/train_configs/training_args_style_sdxl.json +++ b/train_configs/training_args_style_sdxl.json @@ -3,26 +3,18 @@ "sd_model_version": "sdxl", "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip", "concept_mode": "style", - "sample_imgs_lora_scale": 0.7, + "sample_imgs_lora_scale": 0.8, "seed": 1, "resolution": 512, "train_batch_size": 4, "n_sample_imgs": 6, - "max_train_steps": 400, + "max_train_steps": 360, "token_warmup_steps": 0, - "checkpointing_steps": 200, + "checkpointing_steps": 360, "ti_lr": 0.001, - "ti_weight_decay": 0.0005, - - "remove_ti_token_from_prompts": false, - "text_encoder_lora_optimizer": null, - "text_encoder_lora_lr": 1.0e-4, - "text_encoder_lora_weight_decay": 1e-5, - "text_encoder_lora_rank": 12, "unet_lr": 0.001, "lora_rank": 16, "use_dora": false, - "caption_model": "blip", "debug": true } \ No newline at end of file diff --git a/trainer/config.py b/trainer/config.py index e55b03e..30a83dd 100644 --- a/trainer/config.py +++ b/trainer/config.py @@ -56,7 +56,7 @@ class TrainingConfig(BaseModel): 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 = 0.0002 + unet_lr: float = 0.001 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 @@ -68,7 +68,7 @@ class TrainingConfig(BaseModel): 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 + token_attention_loss_w: float = 2e-7 cond_reg_w: float = 0.0e-5 tok_cond_reg_w: float = 0.0e-5 tok_cov_reg_w: float = 500. # regularizes the token covariance matrix wrt pretrained, normal tokens diff --git a/trainer/preprocess.py b/trainer/preprocess.py index c517779..5026d4f 100755 --- a/trainer/preprocess.py +++ b/trainer/preprocess.py @@ -819,16 +819,6 @@ def load_and_save_masks_and_captions( captions = captions + captions - aug_imgs, aug_caps = [],[] - # if we still have a small amount of imgs, do some basic augmentation: - while len(images) + len(aug_imgs) < augment_imgs_up_to_n: - print(f"Adding augmented version of each training img...") - aug_imgs.extend([augment_image(image) for image in images]) - aug_caps.extend(captions) - - images.extend(aug_imgs) - captions.extend(aug_caps) - print(f"Generating {len(images)} captions using mode: {concept_mode}...") captions = caption_dataset(images, captions, caption_model = caption_model) @@ -844,6 +834,19 @@ def load_and_save_masks_and_captions( if not config.disable_ti: captions, trigger_text, gpt_concept_description = post_process_captions(captions, caption_text, concept_mode, seed, skip_gpt_cleanup=config.skip_gpt_cleanup) + + + aug_imgs, aug_caps = [],[] + # if we still have a small amount of imgs, do some basic augmentation: + while len(images) + len(aug_imgs) < augment_imgs_up_to_n: + print(f"Adding augmented version of each training img...") + aug_imgs.extend([augment_image(image) for image in images]) + aug_caps.extend(captions) + + images.extend(aug_imgs) + captions.extend(aug_caps) + + 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 @@ -855,7 +858,6 @@ def load_and_save_masks_and_captions( temp = config.clipseg_temperature print(f"Generating {len(images)} masks...") - # Make sure we have a bias for the background pixels to never 100% ignore them background_bias = 0.05 if not use_face_detection_instead: diff --git a/trainer/ti_cross_attn_loss.py b/trainer/ti_cross_attn_loss.py index 8ea00d9..858a195 100644 --- a/trainer/ti_cross_attn_loss.py +++ b/trainer/ti_cross_attn_loss.py @@ -248,18 +248,19 @@ class DAAMLoss: 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) + if height*width < min_heatmap_pixels: + min_heatmap_pixels = height*width + min_heatmap_shape = height, width # 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) + heatmap = F.interpolate(heatmap.permute(0, 3, 1, 2), size=(min_heatmap_shape[0], min_heatmap_shape[1]), mode='bicubic').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