slow port to a standard config

This commit is contained in:
mayukhdeb
2024-03-15 06:27:18 -07:00
parent 17fcb39190
commit 269eeec1bc
2 changed files with 96 additions and 15 deletions
+54 -15
View File
@@ -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}
+42
View File
@@ -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