From cd9b6629bc5b8f0bc4d073c3d2ac97c83eebe0af Mon Sep 17 00:00:00 2001 From: aiXander Date: Fri, 29 Mar 2024 15:09:10 -0700 Subject: [PATCH] update args --- .gitignore | 2 ++ trainer/config.py | 2 ++ trainer/dataset_and_utils.py | 2 ++ trainer/utils/inference.py | 2 +- trainer_pti.py | 16 ++++++++++------ training_args.json | 4 ++-- 6 files changed, 19 insertions(+), 9 deletions(-) diff --git a/.gitignore b/.gitignore index ee21b64..ba867d0 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,5 @@ tests/ train.py debug/* !debug/*.py + +training_args_x_*.json diff --git a/trainer/config.py b/trainer/config.py index 1ad1562..5cff9c9 100644 --- a/trainer/config.py +++ b/trainer/config.py @@ -23,6 +23,7 @@ class TrainingConfig(BaseModel): ti_weight_decay: float = 3e-4 lora_weight_decay: float = 0.002 l1_penalty: float = 0.1 + noise_offset: float = 0.05 snr_gamma: float = 5.0 lora_rank: int = 12 use_dora: bool = False @@ -55,6 +56,7 @@ class TrainingConfig(BaseModel): lr_num_cycles: int = 1 lr_power: float = 1.0 dataloader_num_workers: int = 0 + training_attributes: dict = {} def save_as_json(self, file_path: str) -> None: with open(file_path, 'w') as f: diff --git a/trainer/dataset_and_utils.py b/trainer/dataset_and_utils.py index 701c624..1125397 100755 --- a/trainer/dataset_and_utils.py +++ b/trainer/dataset_and_utils.py @@ -66,6 +66,8 @@ def plot_loss(losses, save_path='losses.png', window_length=31, polyorder=3): # plt.yscale('log') # Uncomment if log scale is desired plt.xlabel('Step') plt.ylabel('Training Loss') + ymin, ymax = plt.ylim() + plt.ylim(min(ymin, 0), ymax) plt.legend() plt.savefig(save_path) plt.close() diff --git a/trainer/utils/inference.py b/trainer/utils/inference.py index 7968ff9..eb46ee1 100644 --- a/trainer/utils/inference.py +++ b/trainer/utils/inference.py @@ -80,6 +80,6 @@ def render_images(training_pipeline, render_size, lora_path, train_step, seed, i img_grid_path = make_validation_img_grid(lora_path) if not reload_entire_pipeline: # restore the training scheduler - pipeline.scheduler = training_scheduler + training_pipeline.scheduler = training_scheduler return validation_prompts_raw \ No newline at end of file diff --git a/trainer_pti.py b/trainer_pti.py index 1ef0a31..42910c5 100755 --- a/trainer_pti.py +++ b/trainer_pti.py @@ -52,7 +52,11 @@ def main( seed = config.seed, ) - instance_data_dir=os.path.join(input_dir, "captions.csv") + # Update the training attributes with some info from the pre-processing: + config.training_attributes["n_training_imgs"] = n_imgs + config.training_attributes["trigger_text"] = trigger_text + config.training_attributes["segmentation_prompt"] = segmentation_prompt + config.training_attributes["captions"] = captions if config.allow_tf32: torch.backends.cuda.matmul.allow_tf32 = True @@ -161,8 +165,9 @@ def main( #unet.add_adapter(unet_lora_config) unet = get_peft_model(unet, unet_lora_config) + pipe.unet = unet print_trainable_parameters(unet, name = 'unet') - + unet_lora_parameters = list(filter(lambda p: p.requires_grad, unet.parameters())) params_to_optimize = [ @@ -213,7 +218,7 @@ def main( ) train_dataset = PreprocessedDataset( - instance_data_dir, + os.path.join(input_dir, "captions.csv"), tokenizer_one, tokenizer_two, vae, @@ -349,10 +354,9 @@ def main( # Sample noise that we'll add to the latents: noise = torch.randn_like(vae_latent) - noise_offset = 0.05 # TODO, turn this into an input arg and do a grid search - if noise_offset > 0.0: + if config.noise_offset > 0.0: # https://www.crosslabs.org//blog/diffusion-with-offset-noise - noise += noise_offset * torch.randn( + noise += config.noise_offset * torch.randn( (noise.shape[0], noise.shape[1], 1, 1), device=noise.device) bsz = vae_latent.shape[0] diff --git a/training_args.json b/training_args.json index 35039db..dd6da12 100644 --- a/training_args.json +++ b/training_args.json @@ -3,7 +3,7 @@ "name": "unnamed", "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip", "concept_mode": "style", - "sd_model_version": "sd15", + "sd_model_version": "sdxl", "seed": 0, "resolution": 960, "train_batch_size": 4, @@ -17,7 +17,7 @@ "l1_penalty": 0.05, "snr_gamma": 5.0, "lora_rank": 12, - "use_dora": false, + "use_dora": true, "caption_prefix": "in the style of TOK, ", "caption_model": "blip", "left_right_flip_augmentation": true,