small tweaks

This commit is contained in:
aiXander
2024-03-30 21:17:09 -07:00
parent 4e97b176a1
commit 67445d3963
4 changed files with 18 additions and 32 deletions
+1 -1
View File
@@ -196,7 +196,7 @@ def clipseg_mask_generator(
if isinstance(target_prompts, str):
print(
f'Warning: only one target prompt "{target_prompts}" was given, so it will be used for all images'
f'Using "{target_prompts}" as CLIP-segmentation prompt for all images.'
)
target_prompts = [target_prompts] * len(images)
+15 -22
View File
@@ -1,28 +1,21 @@
from ..config import TrainingConfig
def modify_args_based_on_concept_mode(
concept_mode: str,
left_right_flip_augmentation: bool,
mask_target_prompts: str,
clipseg_temperature: float,
l1_penalty: float
):
if concept_mode == "face":
def modify_args_based_on_concept_mode(config: TrainingConfig):
if config.concept_mode == "face":
print(f"Face mode is active ----> disabling left-right flips and setting mask_target_prompts to 'face'.")
left_right_flip_augmentation = False # always disable lr flips for face mode!
mask_target_prompts = "face"
clipseg_temperature = 0.4
config.left_right_flip_augmentation = False # always disable lr flips for face mode!
config.mask_target_prompts = "face"
config.clipseg_temperature = 0.4
if concept_mode == "concept": # gracefully catch any old versions of concept_mode
concept_mode = "object"
if config.concept_mode == "concept": # gracefully catch any old versions of concept_mode
config.concept_mode = "object"
if concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
l1_penalty = 0.05
if config.concept_mode == "style": # for styles you usually want the LoRA matrices to absorb a lot (instead of just the token embedding)
config.l1_penalty = 0.05
return (
concept_mode,
left_right_flip_augmentation,
mask_target_prompts,
clipseg_temperature,
l1_penalty
)
if config.use_dora:
print(f"Disabling L1 penalty and LoRA weight decay for DORA training.")
config.l1_penalty = 0.0
config.lora_weight_decay = 0.0
return config
+1 -8
View File
@@ -70,14 +70,7 @@ def main(
seed_everything(config.seed)
pick_best_gpu_id()
(config.concept_mode, config.left_right_flip_augmentation, config.mask_target_prompts, config.clipseg_temperature, config.l1_penalty
) = modify_args_based_on_concept_mode(
concept_mode=config.concept_mode,
left_right_flip_augmentation=config.left_right_flip_augmentation,
mask_target_prompts=config.mask_target_prompts,
clipseg_temperature=config.clipseg_temperature,
l1_penalty=config.l1_penalty
)
config = modify_args_based_on_concept_mode(config)
input_dir, n_imgs, trigger_text, segmentation_prompt, captions = preprocess(
working_directory=config.output_dir,
+1 -1
View File
@@ -5,7 +5,7 @@
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/clipx_tiny.zip",
"concept_mode": "style",
"caption_prefix": "in the style of TOK, ",
"seed": 0,
"seed": 1,
"resolution": 1024,
"train_batch_size": 4,
"max_train_steps": 500,