init PreprocessedDataset

This commit is contained in:
mayukhdeb
2024-06-20 00:27:01 -07:00
parent cacd6f2201
commit 2c2bfcda13
+100 -15
View File
@@ -11,7 +11,8 @@ todos:
[] - [optional] init lora params for text encoders
[x] - get textual inversion params and it's corresponding optimizer
[x] - either do full finetuning of transformer or init lora params
[] - init PreprocessedDataset object
[x] - init optimizer for transformer trainable parameters
[x] - init PreprocessedDataset object
[] - init OptimizerCollection containing all optimizers
[] - [debug step] visualize a random token embedding
[] - do training. Save checkpoint during and after training
@@ -41,6 +42,14 @@ from trainer.optimizer import (
)
from trainer.optimizer import count_trainable_params
from peft import LoraConfig, get_peft_model
from typing import Iterable
import prodigyopt
import torch
from trainer.preprocess import preprocess
from trainer.dataset import PreprocessedDataset
import argparse
from trainer.config import TrainingConfig
def load_sd3_tokenizers():
# Load the tokenizers
@@ -156,7 +165,37 @@ def freeze_all_gradients(models: list):
for model in models:
model.requires_grad_(False)
def main():
def get_transformer_optimizer(
prodigy_d_coef: float,
prodigy_growth_factor: float,
lora_weight_decay: float,
use_dora: bool,
transformer_trainable_params: Iterable,
optimizer_name="prodigy"
):
if optimizer_name == "adamw":
optimizer = torch.optim.AdamW(transformer_trainable_params, lr = 1e-4, weight_decay=lora_weight_decay if not use_dora else 0.0)
elif optimizer_name == "prodigy":
# Note: the specific settings of Prodigy seem to matter A LOT
optimizer = prodigyopt.Prodigy(
transformer_trainable_params,
d_coef = prodigy_d_coef,
lr=1.0,
decouple=True,
use_bias_correction=True,
safeguard_warmup=True,
weight_decay=lora_weight_decay if not use_dora else 0.0,
betas=(0.9, 0.99),
growth_rate=prodigy_growth_factor # lower values make the lr go up slower (1.01 is for 1k step runs, 1.02 is for 500 step runs)
)
else:
raise NotImplementedError(f"Invalid optimizer_name for unet: {optimizer_name}")
print(f"Created {optimizer_name} optimizer for transformer!")
return optimizer
def main(config: TrainingConfig):
# 1. Load tokenizers
tokenizer_one, tokenizer_two, tokenizer_three = load_sd3_tokenizers()
@@ -202,17 +241,40 @@ def main():
)
# Experimental TODO: warmup the token embeddings using CLIP-similarity optimization
from trainer.config import TrainingConfig
config= TrainingConfig(
lora_training_urls = "none",
concept_mode = "object",
sd_model_version = "sd3",
training_attributes = {
"gpt_description": "A banana with a face"
},
token_warmup_steps = 0,
is_lora = True,
lora_rank = 4
# config= TrainingConfig(
# lora_training_urls = "none",
# concept_mode = "object",
# sd_model_version = "sd3",
# training_attributes = {
# "gpt_description": "A banana with a face"
# },
# token_warmup_steps = 0,
# is_lora = True,
# lora_rank = 4
# )
"""
override some config params because we're recycling an sdxl config here
"""
config.is_lora = True
config.token_warmup_steps = 0
config.lora_rank = 4
config.sd_model_version = "sd3"
config, input_dir = preprocess(
config,
working_directory=config.output_dir,
concept_mode=config.concept_mode,
input_zip_path=config.lora_training_urls,
caption_text=config.caption_prefix,
mask_target_prompts=config.mask_target_prompts,
target_size=config.resolution,
crop_based_on_salience=config.crop_based_on_salience,
use_face_detection_instead=config.use_face_detection_instead,
left_right_flip_augmentation=config.left_right_flip_augmentation,
augment_imgs_up_to_n = config.augment_imgs_up_to_n,
caption_model = config.caption_model,
seed = config.seed,
)
embedding_handler.make_embeddings_trainable()
@@ -254,7 +316,30 @@ def main():
print(f"config.is_lora: {config.is_lora} params: {count_trainable_params(transformer)}")
optimizer_transformer = get_transformer_optimizer(
prodigy_d_coef=config.prodigy_d_coef,
prodigy_growth_factor=config.unet_prodigy_growth_factor,
lora_weight_decay=config.lora_weight_decay,
use_dora=config.use_dora,
transformer_trainable_params=transformer_trainable_params,
optimizer_name=config.unet_optimizer_type
)
print(f"Optimizer: {optimizer_transformer}")
train_dataset = PreprocessedDataset(
input_dir,
pipeline,
vae.float(),
size = config.train_img_size,
do_cache=config.do_cache,
substitute_caption_map=config.token_dict,
aspect_ratio_bucketing=config.aspect_ratio_bucketing,
train_batch_size=config.train_batch_size
)
print(f"train_dataset contains: {len(train_dataset)} samples")
if __name__ == "__main__":
main()
parser = argparse.ArgumentParser(description='Train a concept')
parser.add_argument('config_filename', type=str, help='Input JSON configuration file')
args = parser.parse_args()
config = TrainingConfig.from_json(file_path=args.config_filename)
main(config=config)