Files
edenartlab-sd-lora-trainer/scripts/create_hyperparam_sweep.py
T
2024-04-29 20:12:08 +02:00

153 lines
5.8 KiB
Python

"""
Faces:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_2.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_5.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/steel.zip
Objects:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_all.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/koji_color.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/plantoid_imgs.zip
Styles:
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/does.zip
https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_200.zip
"""
import random, os, ast, json, shutil
from itertools import product
import time
from tqdm import tqdm
random.seed(int(1000*time.time()))
def hamming_distance(dict1, dict2):
distance = 0
for key in dict1.keys():
if dict1[key] != dict2.get(key, None):
distance += 1
return distance
#######################################################################################
# Setup the base experiment config:
exp_name = "sd15_face_sweep"
caption_prefix = ""
mask_target_prompts = ""
n_exp = 200 # how many random experiment settings to generate
min_hamming_distance = 3 # min_n_params that have to be different from any previous experiment to be scheduled
output_sh_path = f"gridsearch_configs/{exp_name}.sh"
# Define training hyperparameters and their possible values
# The params are sampled stochastically, so if you want to use a specific value more often, just put it in multiple times
hyperparameters = {
"output_dir": [f"lora_models/{exp_name}"],
"sd_model_version": ["sd15"],
"lora_training_urls": [
"/home/rednax/Documents/datasets/people/xander"
],
"concept_mode": ['face'],
"seed": [0],
"resolution": [512,640,768],
"train_batch_size": [4],
"n_sample_imgs": [6],
"max_train_steps": [400,800],
"checkpointing_steps": [100],
"gradient_accumulation_steps": [1],
"n_tokens": [1, 2],
"ti_lr": [0.001],
"ti_weight_decay": [0.0005],
"l1_penalty": [0.0],
"token_warmup_steps": [0,60],
"tok_cov_reg_w": [0, 2000],
"cond_reg_w": [0.01e-5, 2.5e-5],
"tok_cond_reg_w": [0.01e-5, 2.5e-5],
"unet_prodigy_growth_factor": [1.05],
"unet_lr": [1.0e-4, 3e-4, 1e-3],
"prodigy_d_coef": [1.0],
"lora_weight_decay": [0.001],
"lora_rank": [6, 12, 24, 48],
"use_dora": ['false', 'true'],
"text_encoder_lora_optimizer": [None],
"text_encoder_lora_lr": [0.0e-4],
"snr_gamma": [5.0],
"caption_model": ["blip"],
"augment_imgs_up_to_n": [20,40],
"verbose": ['true'],
"debug": ['true']
}
#######################################################################################
# Create a set to hold the combinations that have already been run
scheduled_experiments = set()
# if config_output_dir exists, remove it:
config_output_dir = f"gridsearch_configs/{exp_name}"
shutil.rmtree(config_output_dir, ignore_errors=True)
os.makedirs(config_output_dir, exist_ok=True)
# Open the shell script file
try_sampling_n_times = 200
for exp_index in tqdm(range(n_exp)): # number of combinations you want to generate
resamples, combination = 0, None
while resamples < try_sampling_n_times:
experiment_settings = {name: random.choice(values) for name, values in hyperparameters.items()}
resamples += 1
min_distance = float('inf')
for str_experiment_settings in scheduled_experiments:
existing_experiment_settings = dict(sorted(ast.literal_eval(str_experiment_settings)))
distance = hamming_distance(experiment_settings, existing_experiment_settings)
min_distance = min(min_distance, distance)
if min_distance >= min_hamming_distance:
str_experiment_settings = str(sorted(experiment_settings.items()))
scheduled_experiments.add(str_experiment_settings)
# Save the experiment to a JSON file
config_filename = f"{config_output_dir}/{exp_name}_{exp_index:03d}.json"
dirname = os.path.dirname(config_filename)
os.makedirs(dirname, exist_ok=True)
# Make some final adjustments to the experiment settings before saving to disk:
experiment_settings["output_dir"] = f'{experiment_settings["output_dir"]}__{exp_index:03d}'
with open(config_filename, "w") as f:
json.dump(experiment_settings, f, indent=4)
break
if resamples >= try_sampling_n_times:
print(f"\nCould not find a new experiment_setting after random sampling {try_sampling_n_times} times, dumping all experiment_settings to .json files")
break
print(f"\n\n---> Saved {len(scheduled_experiments)} experiment configurations to {config_output_dir}")
def generate_sh_script(folder_path, output_sh_path):
# Get a list of JSON files in the folder, sorted alphabetically
json_files = sorted([f for f in os.listdir(folder_path) if f.endswith('.json')])
# Open the output .sh file for writing
with open(output_sh_path, 'w') as sh_file:
# Write the shebang line for a bash script
sh_file.write("#!/bin/bash\n\n")
# Write a command for each JSON file
for json_file in json_files:
command = f"python main.py {os.path.join(folder_path, json_file)}\n"
sh_file.write(command)
generate_sh_script(config_output_dir, output_sh_path)
print(f"\n---> Saved the executable shell script to {output_sh_path}")