This commit is contained in:
xander
2024-08-05 13:13:38 +02:00
parent ac2dcf787c
commit 2c1bb4a3f9
9 changed files with 51 additions and 149 deletions
+1
View File
@@ -17,6 +17,7 @@ gridsearch*
aesthetic_score_best_model.pth
# experiment folders:
scripts/plots
conditioning_spaces/
training_args_x_*.json
xander_configs/
-87
View File
@@ -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
+6 -6
View File
@@ -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']
+15 -12
View File
@@ -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
+6 -16
View File
@@ -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
}
+3 -11
View File
@@ -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
}
+2 -2
View File
@@ -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
+13 -11
View File
@@ -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:
+5 -4
View File
@@ -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