From 47a0d37b45dea18ef9f2f899e5e93aa2fd0dd710 Mon Sep 17 00:00:00 2001 From: xander Date: Sat, 3 Aug 2024 03:49:00 +0200 Subject: [PATCH] tweak params --- main.py | 3 +-- scripts/create_hyperparam_sweep.py | 4 +--- train_configs/test.json | 10 +++++----- trainer/loss.py | 4 ++-- trainer/preprocess.py | 9 +++++++-- 5 files changed, 16 insertions(+), 14 deletions(-) diff --git a/main.py b/main.py index 16180e4..0f5d9db 100755 --- a/main.py +++ b/main.py @@ -327,8 +327,7 @@ def train(config: TrainingConfig): token_attention_loss = compute_token_attention_loss(pipe, embedding_handler, captions, mask, daam_loss) losses['token_attention_loss'].append(token_attention_loss.item()) - #loss = loss + 0.0000002 * token_attention_loss - loss = loss + 0.000002 * token_attention_loss + loss = loss + 0.0000002 * token_attention_loss if global_step%40 == 0: img_ratio = config.train_img_size[0] / config.train_img_size[1] diff --git a/scripts/create_hyperparam_sweep.py b/scripts/create_hyperparam_sweep.py index 0e2d1f3..1313a01 100644 --- a/scripts/create_hyperparam_sweep.py +++ b/scripts/create_hyperparam_sweep.py @@ -50,9 +50,7 @@ hyperparameters = { "output_dir": [f"lora_models/{exp_name}"], "sd_model_version": ["sdxl"], "lora_training_urls": [ - "/home/rednax/SSD2TB/Github_repos/Eden/images/beeple_large", - "/home/rednax/SSD2TB/Github_repos/Eden/images/beeple" - + "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx.zip", ], "concept_mode": ['style'], "sample_imgs_lora_scale": [0.8], diff --git a/train_configs/test.json b/train_configs/test.json index d7549a8..e4cc955 100644 --- a/train_configs/test.json +++ b/train_configs/test.json @@ -1,8 +1,8 @@ { "name": "xander_test", "sd_model_version": "sdxl", - "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/gene.zip", - "concept_mode": "face", + "lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx.zip", + "concept_mode": "style", "seed": 0, "resolution": 512, "train_batch_size": 4, @@ -15,14 +15,14 @@ "ti_weight_decay": 0.0005, "disable_ti": false, - "n_tokens": 3, + "n_tokens": 2, "text_encoder_lora_optimizer": null, "text_encoder_lora_lr": 1.0e-4, "text_encoder_lora_weight_decay": 1e-5, "text_encoder_lora_rank": 12, - "sample_imgs_lora_scale": 1.0, - "unet_lr": 0.00, + "sample_imgs_lora_scale": 0.8, + "unet_lr": 0.0002, "lora_rank": 16, "use_dora": false, "caption_model": "florence", diff --git a/trainer/loss.py b/trainer/loss.py index 313ee5b..6492ea2 100644 --- a/trainer/loss.py +++ b/trainer/loss.py @@ -54,11 +54,11 @@ def compute_token_attention_loss(pipe, embedding_handler, ti_masks = torch.stack(ti_masks) # Penalize large, positive mean token attentions: - reg_loss_0 = 5*torch.stack(att_L2_losses).mean() + reg_loss_0 = 10*torch.stack(att_L2_losses).mean() # Where the segmentation mask is one, we want to avoid very large attention scores: threshold = 0.0 - reg_loss_1 = (torch.relu(ti_heatmaps*ti_masks - threshold)**2).mean() + reg_loss_1 = 0.5*(torch.relu(ti_heatmaps*ti_masks - threshold)**2).mean() # Where the segmentation mask is zero, we want to avoid somewhat large attention scores: threshold = -15.0 diff --git a/trainer/preprocess.py b/trainer/preprocess.py index adeda4b..e3357d4 100755 --- a/trainer/preprocess.py +++ b/trainer/preprocess.py @@ -557,7 +557,7 @@ def fixed_get_imports(filename: str | os.PathLike) -> list[str]: return imports from transformers import AutoProcessor, AutoModelForCausalLM - +@torch.no_grad() def florence_caption_dataset(images, captions): with patch("transformers.dynamic_module_utils.get_imports", fixed_get_imports): #workaround for unnecessary flash_attn requirement @@ -574,7 +574,7 @@ def florence_caption_dataset(images, captions): input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=1024, - num_beams=random.choice(2,3,4) + num_beams=random.choice([2,3,4]) ) generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0] @@ -582,6 +582,11 @@ def florence_caption_dataset(images, captions): caption = parsed_answer[prompt] captions[i] = caption.replace("The image shows a ", "A ") + del model + del processor + gc.collect() + torch.cuda.empty_cache() + return captions