diff --git a/preprocess.py b/preprocess.py index 3ac8f78..6e22c4f 100755 --- a/preprocess.py +++ b/preprocess.py @@ -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) diff --git a/trainer/utils/config_modification.py b/trainer/utils/config_modification.py index 77822d8..7dde0e2 100644 --- a/trainer/utils/config_modification.py +++ b/trainer/utils/config_modification.py @@ -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 - ) \ No newline at end of file + 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 \ No newline at end of file diff --git a/trainer_pti.py b/trainer_pti.py index 576180c..81f4aca 100755 --- a/trainer_pti.py +++ b/trainer_pti.py @@ -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, diff --git a/training_args.json b/training_args.json index b72726b..cef2e1b 100644 --- a/training_args.json +++ b/training_args.json @@ -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,