From 269eeec1bc328d6e5fbbbaad55c37bcdf385e769 Mon Sep 17 00:00:00 2001 From: mayukhdeb Date: Fri, 15 Mar 2024 06:27:18 -0700 Subject: [PATCH] slow port to a standard config --- predict.py | 69 ++++++++++++++++++++++++++++++++++++----------- trainer/config.py | 42 +++++++++++++++++++++++++++++ 2 files changed, 96 insertions(+), 15 deletions(-) create mode 100644 trainer/config.py diff --git a/predict.py b/predict.py index 26d7c5b..11cf7ad 100755 --- a/predict.py +++ b/predict.py @@ -186,33 +186,72 @@ class Predictor(BasePredictor): bs=6: 8.0 imgs/s, """ + from trainer.config import TrainingConfig + + config = TrainingConfig( + name=name, + lora_training_urls=lora_training_urls, + concept_mode=concept_mode, + sd_model_version=sd_model_version, + seed=seed, + resolution=resolution, + train_batch_size=train_batch_size, + num_train_epochs=num_train_epochs, + max_train_steps=max_train_steps, + checkpointing_steps=checkpointing_steps, + gradient_accumulation_steps=gradient_accumulation_steps, + is_lora=is_lora, + prodigy_d_coef=prodigy_d_coef, + ti_lr=ti_lr, + ti_weight_decay=ti_weight_decay, + lora_weight_decay=lora_weight_decay, + l1_penalty=l1_penalty, + lora_param_scaler=lora_param_scaler, + snr_gamma=snr_gamma, + lora_rank=lora_rank, + caption_prefix=caption_prefix, + caption_model=caption_model, + left_right_flip_augmentation=left_right_flip_augmentation, + augment_imgs_up_to_n=augment_imgs_up_to_n, + n_tokens=n_tokens, + mask_target_prompts=mask_target_prompts, + crop_based_on_salience=crop_based_on_salience, + use_face_detection_instead=use_face_detection_instead, + clipseg_temperature=clipseg_temperature, + verbose=verbose, + run_name=run_name, + debug=debug, + hard_pivot=hard_pivot, + off_ratio_power=off_ratio_power, + ) + start_time = time.time() out_root_dir = "lora_models" - if seed is None: - seed = np.random.randint(0, 2**32 - 1) + if config.seed is None: + config.seed = np.random.randint(0, 2**32 - 1) # Try to make the training reproducible: - random.seed(seed) - np.random.seed(seed) - torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) + random.seed(config.seed) + np.random.seed(config.seed) + torch.manual_seed(config.seed) + torch.cuda.manual_seed_all(config.seed) - if concept_mode == "face": - left_right_flip_augmentation = False # always disable lr flips for face mode! - mask_target_prompts = "face" + if config.concept_mode == "face": + config.left_right_flip_augmentation = False # always disable lr flips for face mode! + config.mask_target_prompts = "face" 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 - print(f"cog:predict:train_lora:{concept_mode}") + print(f"cog:predict:train_lora:{config.concept_mode}") if not debug: - yield CogOutput(name=name, progress=0.0) + yield CogOutput(name=config.name, progress=0.0) # Initialize pretrained_model dictionary pretrained_model = {"version": sd_model_version} diff --git a/trainer/config.py b/trainer/config.py new file mode 100644 index 0000000..dbe2a3c --- /dev/null +++ b/trainer/config.py @@ -0,0 +1,42 @@ +from typing import Union +from pydantic import BaseModel +import random +import json +from typing import Literal +import torch + +class TrainingConfig(BaseModel, extra = "forbid"): + name: str = "unnamed" + lora_training_urls: str + concept_mode: Literal["face", "style", "object"] + sd_model_version: Literal["sdxl", "sd15"] + seed: Union[int, None] = None + resolution: int = 960 + train_batch_size: int = 4 + num_train_epochs: int = 10_000 + max_train_steps: int = 600 + checkpointing_steps: int + gradient_accumulation_steps: int = 1 + is_lora: bool = True + prodigy_d_coef: float = 0.5 + ti_lr: float = 1e-3 + ti_weight_decay: float = 3e-4 + lora_weight_decay: float = 0.002 + l1_penalty: float = 0.1 + lora_param_scaler: float = 0.5 + snr_gamma: float = 5.0 + lora_rank: int = 12 + caption_prefix: str = "" + caption_model: Literal["gpt4-v", "blip"] = "blip" + left_right_flip_augmentation: bool = True + augment_imgs_up_to_n: int = 20 + n_tokens: int = 2 + mask_target_prompts: Union[None, str] + crop_based_on_salience: bool = True + use_face_detection_instead: bool = False + clipseg_temperature: float = 0.7 + verbose: bool = False + run_name: str = "default_run_name" + debug: bool = False + hard_pivot: bool = False + off_ratio_power: float = False \ No newline at end of file