first pass to try and fix ComfyUI loading
This commit is contained in:
@@ -302,7 +302,9 @@ def train(
|
||||
captions, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch()
|
||||
|
||||
captions = list(captions)
|
||||
|
||||
vae_latent = vae_latent.to(pipe.device).to(weight_dtype)
|
||||
mask = mask.to(pipe.device).to(weight_dtype)
|
||||
|
||||
prompt_embeds, pooled_prompt_embeds, add_time_ids = get_conditioning_signals(
|
||||
config, pipe, captions
|
||||
)
|
||||
@@ -423,7 +425,7 @@ def train(
|
||||
is_lora=config.is_lora,
|
||||
unet_lora_parameters=unet_lora_parameters,
|
||||
unet_param_to_optimize=unet_param_to_optimize,
|
||||
name=name
|
||||
name=config.name
|
||||
)
|
||||
last_save_step = global_step
|
||||
|
||||
@@ -486,9 +488,6 @@ def train(
|
||||
|
||||
if not os.path.exists(output_save_dir):
|
||||
os.makedirs(output_save_dir, exist_ok=True)
|
||||
config.save_as_json(
|
||||
os.path.join(output_save_dir, "training_args.json")
|
||||
)
|
||||
save_lora(
|
||||
output_dir=output_save_dir,
|
||||
global_step=global_step,
|
||||
@@ -499,7 +498,7 @@ def train(
|
||||
is_lora=config.is_lora,
|
||||
unet_lora_parameters=unet_lora_parameters,
|
||||
unet_param_to_optimize=unet_param_to_optimize,
|
||||
name=name
|
||||
name=config.name
|
||||
)
|
||||
validation_prompts = render_images(pipe, config.validation_img_size, output_save_dir, global_step,
|
||||
config.seed,
|
||||
@@ -529,11 +528,7 @@ def train(
|
||||
del pipe
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
config.save_as_json(
|
||||
os.path.join(output_save_dir, "training_args.json")
|
||||
)
|
||||
|
||||
|
||||
if config.debug:
|
||||
parent_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
# Create a zipfile of all the *.py files in the directory
|
||||
@@ -543,6 +538,7 @@ def train(
|
||||
|
||||
config.job_time = time.time() - config.start_time
|
||||
config.training_attributes["validation_prompts"] = validation_prompts
|
||||
config.save_as_json(os.path.join(output_save_dir, "training_args.json"))
|
||||
|
||||
return config, output_save_dir
|
||||
|
||||
|
||||
+15
-29
@@ -5,13 +5,14 @@ import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
import os, json, random, time, sys
|
||||
|
||||
sys.path.append('.')
|
||||
sys.path.append('..')
|
||||
from trainer.models import load_models, pretrained_models
|
||||
from trainer.lora import patch_pipe_with_lora
|
||||
from trainer.utils.val_prompts import val_prompts
|
||||
from trainer.utils.io import make_validation_img_grid
|
||||
from trainer.utils.utils import seed_everything, pick_best_gpu_id
|
||||
from trainer.utils.inference import encode_prompt_advanced
|
||||
from trainer.inference import encode_prompt_advanced
|
||||
|
||||
def load_model(pretrained_model):
|
||||
if pretrained_model['version'] == "sd15":
|
||||
@@ -29,16 +30,16 @@ def load_model(pretrained_model):
|
||||
if __name__ == "__main__":
|
||||
|
||||
pretrained_model = pretrained_models['sdxl']
|
||||
lora_path = 'lora_models/xander--07_19-19-35-sdxl_face_dora/checkpoints/checkpoint-360'
|
||||
lora_scales = np.linspace(0.4, 0.6, 3)
|
||||
render_size = (1024, 256+1024) # H,W
|
||||
lora_path = 'lora_models/lizzo--09_02-04-26-sdxl_object_lora_640_0.1_blip/checkpoints/checkpoint-500'
|
||||
lora_scales = np.linspace(0.5, 0.7, 3)
|
||||
render_size = (768, 768) # H,W
|
||||
n_imgs = 10
|
||||
n_loops = 4
|
||||
|
||||
n_steps = 35
|
||||
guidance_scale = 7.5
|
||||
seed = 3
|
||||
use_lightning = 1
|
||||
seed = 2
|
||||
use_lightning = 0
|
||||
|
||||
#####################################################################################
|
||||
|
||||
@@ -70,27 +71,11 @@ if __name__ == "__main__":
|
||||
validation_prompts_raw = random.choices(val_prompts['object'], k=n_imgs)
|
||||
|
||||
validation_prompts_raw = [
|
||||
'an elderly <concept>, old man',
|
||||
'a photo of <concept> as a young boy',
|
||||
'a painting of elderly <concept>, impressionism, stunning composition',
|
||||
'a painting of <concept> as a young kid, playing in the garden',
|
||||
'A digital illustration of <concept> as a child, exploring a mystical forest, vibrant colors',
|
||||
'An old-fashioned portrait of elderly <concept>, sitting in a cozy armchair, reading a book',
|
||||
'A watercolor painting of young <concept>, running through a meadow with a kite, bright and joyful',
|
||||
'A sketch of <concept> as a teenager, gazing at the stars, dreamy and contemplative',
|
||||
'A vintage photograph of elderly <concept>, wearing a classic hat, exuding wisdom and elegance',
|
||||
'A cartoon drawing of young <concept>, having a playful snowball fight, full of energy and laughter',
|
||||
'A surrealist painting of <concept> as a child, riding a fantastical creature, imaginative and whimsical',
|
||||
'A black and white photo of elderly <concept>, sitting on a park bench, reflecting on lifes journey',
|
||||
'A futuristic hologram of elderly <concept>, floating in a space station, surrounded by stars and galaxies'
|
||||
'A clay animation of baby <concept>, crawling in a magical garden with talking flowers and dancing insects',
|
||||
'An abstract painting of ancient <concept>, merging with the roots of an ancient tree, symbolizing wisdom and eternity',
|
||||
'A neon-lit digital art piece of toddler <concept>, playing with holographic toys in a cyberpunk playground',
|
||||
'A mosaic of centenarian <concept>, composed of thousands of tiny images from their life, telling a story of a century',
|
||||
'A pop art portrait of young <concept>, riding a colorful unicorn, surrounded by rainbows and candy clouds',
|
||||
'A steampunk illustration of elderly <concept>, inventing a time machine, surrounded by gears and steam',
|
||||
'A fantasy drawing of newborn <concept>, cradled in the arms of a gentle giant, in a land of giants and mythical creatures'
|
||||
|
||||
"actress in TOK at a gala",
|
||||
"TOK, facy dinner party",
|
||||
"woman in TOK and jacket",
|
||||
"smiling woman in TOK at a party",
|
||||
"jennifer jones on TOK at music awards"
|
||||
]
|
||||
|
||||
negative_prompt = "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft"
|
||||
@@ -103,8 +88,9 @@ if __name__ == "__main__":
|
||||
for jj in range(n_loops):
|
||||
for i in range(len(validation_prompts_raw)):
|
||||
for lora_scale in lora_scales:
|
||||
pipe = load_model(pretrained_model)
|
||||
pipe.unet = PeftModel.from_pretrained(model = pipe.unet, model_id = lora_path, adapter_name = 'eden_lora')
|
||||
seed += 1
|
||||
#pipe = load_model(pretrained_model)
|
||||
#pipe.unet = PeftModel.from_pretrained(model = pipe.unet, model_id = lora_path, adapter_name = 'eden_lora')
|
||||
pipe = patch_pipe_with_lora(pipe, lora_path, lora_scale=lora_scale)
|
||||
generator = torch.Generator(device='cuda').manual_seed(seed)
|
||||
|
||||
|
||||
+5
-1
@@ -46,7 +46,7 @@ class TrainingConfig(BaseModel):
|
||||
clipseg_temperature: float = 0.6
|
||||
n_sample_imgs: int = 4
|
||||
verbose: bool = False
|
||||
name: str = "unnamed"
|
||||
name: str = None
|
||||
output_dir: str = "lora_models/unnamed"
|
||||
debug: bool = False
|
||||
hard_pivot: bool = False
|
||||
@@ -80,6 +80,10 @@ class TrainingConfig(BaseModel):
|
||||
# add some metrics to the foldername:
|
||||
lora_str = "dora" if self.use_dora else "lora"
|
||||
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
|
||||
|
||||
if not self.name:
|
||||
self.name = f"{os.path.basename(self.output_dir)}_{self.concept_mode}_{lora_str}_{self.sd_model_version}_{timestamp_short}"
|
||||
|
||||
self.output_dir = self.output_dir + f"--{timestamp_short}-{self.sd_model_version}_{self.concept_mode}_{lora_str}_{self.resolution}_{self.prodigy_d_coef}_{self.caption_model}"
|
||||
os.makedirs(self.output_dir, exist_ok=True)
|
||||
|
||||
|
||||
@@ -32,8 +32,6 @@ class PreprocessedDataset(Dataset):
|
||||
data_dir: str,
|
||||
pipe,
|
||||
vae_encoder,
|
||||
text_encoder_1=None,
|
||||
text_encoder_2=None,
|
||||
do_cache: bool = False,
|
||||
size: List[int] = [512, 512],
|
||||
text_dropout: float = 0.0,
|
||||
@@ -59,10 +57,6 @@ class PreprocessedDataset(Dataset):
|
||||
self.mask_path = self.data["mask_path"]
|
||||
|
||||
self.pipe = pipe
|
||||
|
||||
self.text_encoder_1 = text_encoder_1
|
||||
self.text_encoder_2 = text_encoder_2
|
||||
|
||||
self.vae_encoder = vae_encoder
|
||||
self.vae_scaling_factor = self.vae_encoder.config.scaling_factor
|
||||
self.text_dropout = text_dropout
|
||||
@@ -93,7 +87,6 @@ class PreprocessedDataset(Dataset):
|
||||
for idx in range(len(self.data)):
|
||||
aspect_ratios[idx] = Image.open(os.path.join(self.data_dir, self.image_path[idx])).size
|
||||
|
||||
print(aspect_ratios)
|
||||
self.bucket_manager = BucketManager(
|
||||
aspect_ratios = aspect_ratios,
|
||||
bsz = train_batch_size,
|
||||
@@ -108,7 +101,6 @@ class PreprocessedDataset(Dataset):
|
||||
indices, resolution = self.bucket_manager.get_batch()
|
||||
|
||||
print(f"Got bucket batch: {indices}, resolution: {resolution}")
|
||||
|
||||
tok1, tok2, vae_latents, masks = [], [], [], []
|
||||
|
||||
for idx in indices:
|
||||
|
||||
+18
-16
@@ -1,10 +1,9 @@
|
||||
import os, json
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
from safetensors.torch import load_file, save_file
|
||||
from typing import Dict
|
||||
from peft import PeftModel
|
||||
from trainer.embedding_handler import TokenEmbeddingsHandler
|
||||
from safetensors.torch import save_file
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
|
||||
def patch_pipe_with_lora(pipe, lora_path, lora_scale = 1.0):
|
||||
@@ -74,13 +73,19 @@ def save_lora(
|
||||
name: str = None
|
||||
):
|
||||
"""
|
||||
Save the LORA model to output_dir
|
||||
Save the model + embeddings to output_dir
|
||||
"""
|
||||
print(f"Saving checkpoint at step.. {global_step}")
|
||||
|
||||
# Make sure all weird delimiter characters are removed from concept_name before using it as a filepath:
|
||||
name = name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_")
|
||||
|
||||
embedding_handler.save_embeddings(f"{output_dir}/{name}_embeddings.safetensors")
|
||||
|
||||
print("A")
|
||||
with open(f"{output_dir}/special_params.json", "w") as f:
|
||||
json.dump(token_dict, f)
|
||||
print("A")
|
||||
|
||||
if not is_lora:
|
||||
lora_tensors = {
|
||||
name: param
|
||||
@@ -89,11 +94,9 @@ def save_lora(
|
||||
}
|
||||
save_file(lora_tensors, f"{output_dir}/unet.safetensors",)
|
||||
elif len(unet_lora_parameters) > 0:
|
||||
unet.save_pretrained(save_directory = output_dir)
|
||||
#unet.save_pretrained(save_directory = output_dir)
|
||||
|
||||
lora_tensors = get_peft_model_state_dict(unet)
|
||||
|
||||
if 1:
|
||||
lora_tensors = get_peft_model_state_dict(unet)
|
||||
unet_lora_layers_to_save = convert_state_dict_to_diffusers(lora_tensors)
|
||||
StableDiffusionXLPipeline.save_lora_weights(
|
||||
output_dir,
|
||||
@@ -102,14 +105,13 @@ def save_lora(
|
||||
#text_encoder_2_lora_layers=text_encoder_two_lora_layers_to_save,
|
||||
)
|
||||
|
||||
if 1:
|
||||
#lora_tensors = unet_attn_processors_state_dict(unet)
|
||||
save_file(lora_tensors, f"{output_dir}/{name}_lora_orig.safetensors")
|
||||
|
||||
embedding_handler.save_embeddings(f"{output_dir}/{name}_embeddings.safetensors")
|
||||
|
||||
with open(f"{output_dir}/special_params.json", "w") as f:
|
||||
json.dump(token_dict, f)
|
||||
# Convert to WebUI format
|
||||
lora_state_dict = load_file(f"{output_dir}/pytorch_lora_weights.safetensors")
|
||||
peft_state_dict = convert_all_state_dict_to_peft(lora_state_dict)
|
||||
kohya_state_dict = convert_state_dict_to_kohya(peft_state_dict)
|
||||
save_file(kohya_state_dict, f"{output_dir}/{name}.safetensors")
|
||||
else:
|
||||
print("No lora parameters to save.")
|
||||
|
||||
|
||||
######################################################
|
||||
|
||||
@@ -731,7 +731,7 @@ def calculate_new_dimensions(target_size, target_aspect_ratio):
|
||||
new_width = new_width - (new_width % 64)
|
||||
new_height = new_height - (new_height % 64)
|
||||
|
||||
return new_width, new_height
|
||||
return [new_width, new_height]
|
||||
|
||||
|
||||
def load_and_save_masks_and_captions(
|
||||
|
||||
@@ -7,16 +7,16 @@
|
||||
"resolution": 512,
|
||||
"validation_img_size": [1024, 1024],
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"n_sample_imgs": 4,
|
||||
"max_train_steps": 360,
|
||||
"token_warmup_steps": 100,
|
||||
"token_warmup_steps": 40,
|
||||
"checkpointing_steps": 60,
|
||||
"gradient_accumulation_steps": 2,
|
||||
"gradient_accumulation_steps": 1,
|
||||
"n_tokens": 3,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0002,
|
||||
"off_ratio_power": 0.05,
|
||||
"prodigy_d_coef": 0.0,
|
||||
"prodigy_d_coef": 0.5,
|
||||
"lora_weight_decay": 0.0001,
|
||||
"l1_penalty": 0.1,
|
||||
"noise_offset": 0.05,
|
||||
@@ -24,7 +24,7 @@
|
||||
"lora_alpha_multiplier": 1.0,
|
||||
"lora_rank": 12,
|
||||
"use_dora": true,
|
||||
"caption_model": "gpt4-v",
|
||||
"caption_model": "blip",
|
||||
"left_right_flip_augmentation": true,
|
||||
"aspect_ratio_bucketing": false,
|
||||
"augment_imgs_up_to_n": 20,
|
||||
@@ -32,6 +32,5 @@
|
||||
"debug": true,
|
||||
"hard_pivot": false,
|
||||
"weight_type": "bf16",
|
||||
"dataloader_num_workers": 0,
|
||||
"name": "unnamed"
|
||||
"dataloader_num_workers": 0
|
||||
}
|
||||
@@ -1,22 +1,22 @@
|
||||
{
|
||||
"output_dir": "lora_models/banny",
|
||||
"sd_model_version": "sdxl",
|
||||
"lora_training_urls": "/home/xander/Downloads/datasets/banny_best_uncaptioned",
|
||||
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/banny_best.zip",
|
||||
"concept_mode": "object",
|
||||
"seed": 1,
|
||||
"resolution": 512,
|
||||
"validation_img_size": [960, 640],
|
||||
"validation_img_size": [1024, 1024],
|
||||
"train_batch_size": 4,
|
||||
"n_sample_imgs": 6,
|
||||
"n_sample_imgs": 4,
|
||||
"max_train_steps": 400,
|
||||
"token_warmup_steps": 100,
|
||||
"token_warmup_steps": 40,
|
||||
"checkpointing_steps": 100,
|
||||
"gradient_accumulation_steps": 1,
|
||||
"n_tokens": 3,
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0002,
|
||||
"off_ratio_power": 0.05,
|
||||
"prodigy_d_coef": 0.1,
|
||||
"prodigy_d_coef": 1.0,
|
||||
"lora_weight_decay": 0.0001,
|
||||
"l1_penalty": 0.5,
|
||||
"noise_offset": 0.05,
|
||||
@@ -32,6 +32,5 @@
|
||||
"debug": true,
|
||||
"hard_pivot": false,
|
||||
"weight_type": "bf16",
|
||||
"dataloader_num_workers": 0,
|
||||
"name": "unnamed"
|
||||
"dataloader_num_workers": 0
|
||||
}
|
||||
@@ -16,7 +16,7 @@
|
||||
"ti_lr": 0.001,
|
||||
"ti_weight_decay": 0.0002,
|
||||
"off_ratio_power": 0.05,
|
||||
"prodigy_d_coef": 0.0,
|
||||
"prodigy_d_coef": 1.0,
|
||||
"lora_weight_decay": 0.0001,
|
||||
"l1_penalty": 0.1,
|
||||
"noise_offset": 0.05,
|
||||
@@ -32,6 +32,5 @@
|
||||
"debug": true,
|
||||
"hard_pivot": false,
|
||||
"weight_type": "bf16",
|
||||
"dataloader_num_workers": 0,
|
||||
"name": "unnamed"
|
||||
"dataloader_num_workers": 0
|
||||
}
|
||||
Reference in New Issue
Block a user