From fc54948efff5e8a2d3d97de103a7ed16885ba4b5 Mon Sep 17 00:00:00 2001 From: mayukhdeb Date: Thu, 11 Jul 2024 12:39:30 -0700 Subject: [PATCH] more tweaks --- generate_sweep_params.py | 7 ++++--- main_sd3.py | 9 +++++---- 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/generate_sweep_params.py b/generate_sweep_params.py index 91137fd..08905ec 100644 --- a/generate_sweep_params.py +++ b/generate_sweep_params.py @@ -52,9 +52,9 @@ default_config = { "resolution": 512, "train_batch_size": 2, "n_sample_imgs": 6, - "max_train_steps": 100, + "max_train_steps": 1000, "token_warmup_steps": 200, - "checkpointing_steps": 1000000000, ## no need to save any checkpoints + "checkpointing_steps": 1000, ## no need to save any checkpoints "gradient_accumulation_steps": 2, "sample_imgs_lora_scale": 0.8, "n_tokens": 2, @@ -68,7 +68,7 @@ default_config = { "lora_rank": 16, "use_dora": False, "caption_model": "blip", - "debug": True + "debug": True, } keys, values = zip(*sweep_params.items()) @@ -88,6 +88,7 @@ for index, c in enumerate(combinations): if key == "train_batch_size": config["gradient_accumulation_steps"] = c[key] / config["train_batch_size"] config["max_train_steps"] = config["max_train_steps"] * config["gradient_accumulation_steps"] + config["checkpointing_steps"] = config["checkpointing_steps"] * config["gradient_accumulation_steps"] else: config[key] = c[key] # print(f"{index} - Setting {key} to {c[key]}") diff --git a/main_sd3.py b/main_sd3.py index e9e725f..3f313f7 100644 --- a/main_sd3.py +++ b/main_sd3.py @@ -841,10 +841,11 @@ def main(config: TrainingConfig, wandb_log = False, output_dir = None): """ Save intermediate checkpoint """ - save_transformer_lora_checkpoint( - transformer=transformer, - folder=os.path.join(checkpoints_folder,f"global_step_{global_step}", f"transformer") - ) + # save_transformer_lora_checkpoint( + # transformer=transformer, + # folder=os.path.join(checkpoints_folder,f"global_step_{global_step}", f"transformer") + # ) + # temporarily commented out since we don't want to save checkpoints and fill up the storage """ Run inference on a few prompts