bugfix
This commit is contained in:
@@ -17,6 +17,7 @@ gridsearch*
|
||||
aesthetic_score_best_model.pth
|
||||
|
||||
# experiment folders:
|
||||
scripts/plots
|
||||
conditioning_spaces/
|
||||
training_args_x_*.json
|
||||
xander_configs/
|
||||
|
||||
@@ -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
|
||||
@@ -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']
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user