tweak params
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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",
|
||||
|
||||
+2
-2
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user