small tweaks
This commit is contained in:
+1
-1
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user