inference updates

This commit is contained in:
aiXander
2024-04-06 00:28:12 -07:00
parent 3870c8ebda
commit ad952988db
8 changed files with 234 additions and 218 deletions
+1 -1
View File
@@ -13,7 +13,7 @@ build:
python_packages:
- "scipy"
- "diffusers"
- "peft==0.9.0"
- "peft==0.10.0"
- "torch==2.0.1"
- "transformers==4.31.0"
- "invisible-watermark==0.2.0"
+23 -56
View File
@@ -1,35 +1,32 @@
from trainer.models import load_models, pretrained_models
from trainer.utils.lora import patch_pipe_with_lora, blend_conditions
from trainer.utils.lora import patch_pipe_with_lora
from trainer.utils.val_prompts import val_prompts
from trainer.utils.prompt import prepare_prompt_for_lora
from trainer.utils.io import make_validation_img_grid
from trainer.dataset_and_utils import pick_best_gpu_id
from trainer.utils.seed import seed_everything
from diffusers import EulerDiscreteScheduler
from trainer.utils.inference import encode_prompt_advanced
import numpy as np
import torch
from huggingface_hub import hf_hub_download
import os, json, random, time
if __name__ == "__main__":
pretrained_model = pretrained_models['sdxl']
lora_path = 'lora_models/banny---sdxl_object_dora/checkpoints/checkpoint-800'
lora_scale = 1.0
modulate_token_strength = True
seed = 1
render_size = (1024, 1024) # H,W
n_imgs = 24
n_steps = 30
lora_path = 'lora_models/plantoid_best--05_21-39-46-sdxl_object_dora/checkpoints/checkpoint-400'
lora_scale = 0.5
render_size = (1024+1024, 1024) # H,W
n_imgs = 30
n_steps = 25
guidance_scale = 8
use_lightning = False
seed = 1
use_lightning = False
#####################################################################################
output_dir = f'test_images/{lora_path.split("/")[-1]}'
output_dir = f'test_images4/{lora_path.split("/")[-1]}'
os.makedirs(output_dir, exist_ok=True)
seed_everything(seed)
@@ -50,15 +47,11 @@ if __name__ == "__main__":
pipe.load_lora_weights(hf_hub_download(repo, ckpt))
pipe.fuse_lora()
n_steps = 8
guidance_scale=1
pipe = patch_pipe_with_lora(pipe, lora_path, lora_scale=lora_scale)
guidance_scale=1.5
with open(os.path.join(lora_path, "training_args.json"), "r") as f:
training_args = json.load(f)
concept_mode = training_args["concept_mode"]
trigger_text = training_args["training_attributes"]["trigger_text"]
segmentation_prompt = training_args["training_attributes"]["segmentation_prompt"]
if concept_mode == "style":
validation_prompts_raw = random.choices(val_prompts['style'], k=n_imgs)
@@ -67,10 +60,8 @@ if __name__ == "__main__":
else:
validation_prompts_raw = random.choices(val_prompts['object'], k=n_imgs)
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path, verbose=1, trigger_text=trigger_text) for prompt in validation_prompts_raw]
pipe.scheduler = EulerDiscreteScheduler.from_config(pipe.scheduler.config, timestep_spacing="trailing")
pipe = patch_pipe_with_lora(pipe, lora_path, lora_scale=lora_scale)
pipe.scheduler = EulerDiscreteScheduler.from_config(pipe.scheduler.config) #, timestep_spacing="trailing")
generator = torch.Generator(device='cuda').manual_seed(seed)
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"
pipeline_args = {
@@ -80,38 +71,14 @@ if __name__ == "__main__":
"width": render_size[1],
}
#cross_attention_kwargs = {"scale": lora_scale}
cross_attention_kwargs = None
kwargs_str = 'none' if cross_attention_kwargs is None else f'kwargs'
for i in range(len(validation_prompts_raw)):
c, uc, pc, puc = encode_prompt_advanced(pipe, lora_path, validation_prompts_raw[i], negative_prompt, lora_scale, guidance_scale)
for i in range(n_imgs):
pipeline_args['prompt_embeds'] = c
pipeline_args['negative_prompt_embeds'] = uc
if pretrained_model['version'] == 'sdxl':
pipeline_args['pooled_prompt_embeds'] = pc
pipeline_args['negative_pooled_prompt_embeds'] = puc
if modulate_token_strength:
embeds = pipe.encode_prompt(
validation_prompts[i],
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt)
zero_prompt = validation_prompts_raw[i].replace('<concept>', segmentation_prompt)
zero_embeds = pipe.encode_prompt(
zero_prompt,
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt)
embeds, token_scale = blend_conditions(zero_embeds, embeds, lora_scale)
c, uc, pc, puc = embeds
pipeline_args['prompt_embeds'] = c
pipeline_args['negative_prompt_embeds'] = uc
if pretrained_model['version'] == 'sdxl':
pipeline_args['pooled_prompt_embeds'] = pc
pipeline_args['negative_pooled_prompt_embeds'] = puc
else:
pipeline_args["prompt"] = validation_prompts[i]
pipeline_args["negative_prompt"] = negative_prompt
token_scale = 1.0
print(f"Rendering test img with prompt: {validation_prompts[i]}")
image = pipe(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
image.save(os.path.join(output_dir, f"img_seed_{seed}_{i}_tok_scale_{token_scale:.2f}_lora_scale_{lora_scale:.2f}_{kwargs_str}_{int(time.time())}.jpg"), format="JPEG", quality=95)
image = pipe(**pipeline_args, generator=generator).images[0]
image.save(os.path.join(output_dir, f"{validation_prompts_raw[i][:40]}_seed_{seed}_{i}_lora_scale_{lora_scale:.2f}_{int(time.time())}.jpg"), format="JPEG", quality=95)
+2 -2
View File
@@ -41,7 +41,7 @@ class TrainingConfig(BaseModel):
mask_target_prompts: Union[None, str] = None
crop_based_on_salience: bool = True
use_face_detection_instead: bool = False
clipseg_temperature: float = 0.7
clipseg_temperature: float = 0.6
n_sample_imgs: int = 4
verbose: bool = False
run_name: str = "default_run_name"
@@ -84,5 +84,5 @@ class TrainingConfig(BaseModel):
lora_str = "dora" if config_data["use_dora"] else "lora"
timestamp_short = datetime.now().strftime("%d_%H-%M-%S")
config_data["output_dir"] = config_data["output_dir"] + f"--{timestamp_short}-{config_data['sd_model_version']}_{config_data['concept_mode']}_{lora_str}"
return cls(**config_data)
+173 -12
View File
@@ -4,13 +4,112 @@ import random
import shutil
import json
import gc
import re
from diffusers import EulerDiscreteScheduler
from .val_prompts import val_prompts
from ..models import load_models
from .lora import patch_pipe_with_lora
from .prompt import prepare_prompt_for_lora
from .io import make_validation_img_grid
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
trigger_text = training_args["training_attributes"]["trigger_text"]
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
lora_name = "concept"
lora_name_encapsulated = "<" + lora_name + ">"
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK"
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
# Fix common mistakes
fix_replacements = {
r",,": ",",
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
r"\s\.": ".",
r"\s,": ","
}
prompt = replace_in_string(prompt, fix_replacements)
if verbose:
print('-------------------------')
print("Adjusted prompt for LORA:")
print(orig_prompt)
print('-- to:')
print(prompt)
print('-------------------------')
return prompt
def get_conditioning_signals(config, pipe, token_indices, text_encoders, weight_dtype):
@@ -76,11 +175,72 @@ def get_conditioning_signals(config, pipe, token_indices, text_encoders, weight_
return prompt_embeds, pooled_prompt_embeds, add_time_ids
def blend_conditions(embeds1, embeds2, lora_scale,
token_scale_power = 0.5, # adjusts the curve of the interpolation
min_token_scale = 0.5, # minimum token scale (corresponds to lora_scale = 0)
token_scale = None,
verbose = True,
):
"""
using lora_scale, apply linear interpolation between two sets of embeddings
"""
c1, uc1, pc1, puc1 = embeds1
c2, uc2, pc2, puc2 = embeds2
if token_scale is None: # compute the token_scale based on lora_scale:
token_scale = lora_scale ** token_scale_power
# rescale the [0,1] range to [min_token_scale, 1] range:
token_scale = min_token_scale + (1 - min_token_scale) * token_scale
if verbose:
print(f"Setting token_scale to {token_scale:.2f} (lora_scale = {lora_scale}, power = {token_scale_power})")
print('-------------------------')
try:
c = (1 - token_scale) * c1 + token_scale * c2
pc = (1 - token_scale) * pc1 + token_scale * pc2
uc = (1 - token_scale) * uc1 + token_scale * uc2
puc = (1 - token_scale) * puc1 + token_scale * puc2
except:
print(f"Error in blending conditions, reverting to c2, uc2, pc2, puc2")
token_scale = 1.0
c = c2
pc = pc2
uc = uc2
puc = puc2
return (c, uc, pc, puc), token_scale
def encode_prompt_advanced(pipe, lora_path, prompt, negative_prompt, lora_scale, guidance_scale):
"""
Helper function to encode the lora_prompt (containing a trained token) and a zero prompt (without the token)
This allows interpolating the strength of the trained token in the final image.
"""
lora_prompt = prepare_prompt_for_lora(prompt, lora_path, verbose=1)
zero_prompt = prompt.replace('<concept>', "")
print(f'Embedding lora prompt: {lora_prompt}')
embeds = pipe.encode_prompt(
lora_prompt,
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt)
print(f"Embedding zero prompt: {zero_prompt}")
zero_embeds = pipe.encode_prompt(
zero_prompt,
do_classifier_free_guidance=guidance_scale > 1,
negative_prompt=negative_prompt)
embeds, token_scale = blend_conditions(zero_embeds, embeds, lora_scale)
return embeds
@torch.no_grad()
def render_images(pipe, render_size, lora_path, train_step, seed, is_lora, pretrained_model, trigger_text: str, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0", verbose: bool = True):
def render_images(pipe, render_size, lora_path, train_step, seed, is_lora, pretrained_model, trigger_text: str, lora_scale = 0.75, n_steps = 25, n_imgs = 4, device = "cuda:0", verbose: bool = True):
training_scheduler = pipe.scheduler
random.seed(seed)
@@ -121,25 +281,26 @@ def render_images(pipe, render_size, lora_path, train_step, seed, is_lora, pretr
pipe.vae = pipe.vae.to(device).to(pipe.unet.dtype)
pipe.scheduler = EulerDiscreteScheduler.from_config(pipe.scheduler.config, timestep_spacing="trailing")
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path, verbose=verbose, trigger_text=trigger_text) for prompt in validation_prompts_raw]
generator = torch.Generator(device=device).manual_seed(0)
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"
pipeline_args = {
"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",
"num_inference_steps": n_steps,
"guidance_scale": 8,
"height": render_size[0],
"width": render_size[1],
}
if is_lora > 0:
cross_attention_kwargs = {"scale": lora_scale}
else:
cross_attention_kwargs = None
for i in range(n_imgs):
pipeline_args["prompt"] = validation_prompts[i]
print(f"Rendering validation img with prompt: {validation_prompts[i]}")
image = pipe(**pipeline_args, generator=generator, cross_attention_kwargs = cross_attention_kwargs).images[0]
print(f"Rendering validation img with prompt: {validation_prompts_raw[i]}")
c, uc, pc, puc = encode_prompt_advanced(pipe, lora_path, validation_prompts_raw[i], negative_prompt, lora_scale, guidance_scale = 8)
pipeline_args['prompt_embeds'] = c
pipeline_args['negative_prompt_embeds'] = uc
if pretrained_model['version'] == 'sdxl':
pipeline_args['pooled_prompt_embeds'] = pc
pipeline_args['negative_pooled_prompt_embeds'] = puc
image = pipe(**pipeline_args, generator=generator).images[0]
image.save(os.path.join(lora_path, f"img_{train_step:04d}_{i}.jpg"), format="JPEG", quality=95)
img_grid_path = make_validation_img_grid(lora_path)
+6 -6
View File
@@ -96,11 +96,11 @@ def merge_datasets(path_A, path_B, out_path, token_names):
def make_validation_img_grid(img_folder):
def make_validation_img_grid(img_folder, rows = 2):
"""
find all the .jpg imgs in img_folder (template = *.jpg)
if >=4 validation imgs, create a 2xn grid of them
if >=4 validation imgs, create a rows x n grid of them
otherwise just return the first validation img
"""
@@ -113,20 +113,20 @@ def make_validation_img_grid(img_folder):
return os.path.join(img_folder, validation_imgs[0])
else:
# If >= 4 validation images, create 2xn grid
n_imgs = len(validation_imgs) // 2 * 2
n_imgs = len(validation_imgs) // rows * rows
imgs = [Image.open(os.path.join(img_folder, img)) for img in validation_imgs[:n_imgs]]
n_cols = int(n_imgs / 2)
n_cols = int(n_imgs / rows)
# Assuming all images are the same size, get dimensions of first image
width, height = imgs[0].size
# Create an empty image with 2x2 grid size
grid_img = Image.new("RGB", (n_cols * width, 2 * height))
grid_img = Image.new("RGB", (n_cols * width, rows * height))
# Paste the images into the grid
for i in range(n_cols):
for j in range(2):
for j in range(rows):
grid_img.paste(imgs.pop(0), (i * width, j * height))
# Save the new image
+26 -37
View File
@@ -16,62 +16,51 @@ from diffusers.utils import (
)
'''
"""
def blend_conditions(embeds1, embeds2, lora_scale,
token_scale_power = 0.5, # adjusts the curve of the interpolation
min_token_scale = 0.5, # minimum token scale (corresponds to lora_scale = 0)
verbose = True,
):
"""
using lora_scale, apply linear interpolation between two sets of embeddings
"""
from peft import get_peft_model
c1, uc1, pc1, puc1 = embeds1
c2, uc2, pc2, puc2 = embeds2
base_model = ... # load the base model, e.g. from transformers
peft_model = PeftMixedModel.from_pretrained(base_model, path_to_adapter1, "adapter1").eval()
peft_model.load_adapter(path_to_adapter2, "adapter2")
peft_model.set_adapter(["adapter1", "adapter2"]) # activate both adapters
peft_model(data) # forward pass using both adapters
token_scale = lora_scale ** token_scale_power
# rescale the [0,1] range to [min_token_scale, 1] range:
token_scale = min_token_scale + (1 - min_token_scale) * token_scale
"""
if verbose:
print(f"Setting token_scale to {token_scale:.2f} (lora_scale = {lora_scale}, power = {token_scale_power})")
print('-------------------------')
try:
c = (1 - token_scale) * c1 + token_scale * c2
pc = (1 - token_scale) * pc1 + token_scale * pc2
uc = (1 - token_scale) * uc1 + token_scale * uc2
puc = (1 - token_scale) * puc1 + token_scale * puc2
except:
print(f"Error in blending conditions, reverting to c2, uc2, pc2, puc2")
token_scale = 1.0
c = c2
pc = pc2
uc = uc2
puc = puc2
return (c, uc, pc, puc), token_scale
def patch_pipe_with_lora(pipe, lora_path, lora_scale = 1.0):
"""
update the pipe with the lora model and the token embeddings
"""
pipe.unet = PeftModel.from_pretrained(pipe.unet, lora_path)
pipe.unet.merge_adapter()
pipe.unet = PeftModel.from_pretrained(model = pipe.unet, model_id = lora_path, adapter_name = 'eden_lora')
#peft_model.load_adapter(path_to_adapter2, "adapter2")
#peft_model.set_adapter(["adapter1", "adapter2"]) # activate both adapters
# First lets see if any lora's are active and unload them:
#pipe.unet.unmerge_adapter()
list_adapters_component_wise = pipe.get_list_adapters()
print(f"list_adapters_component_wise: {list_adapters_component_wise}")
#pipe.load_lora_weights(lora_path, adapter_name = "my_adapter", weight_name="pytorch_lora_weights.safetensors")
#state_dict, network_alphas = pipe.lora_state_dict(lora_path)
#for key in state_dict.keys():
# print(f"{key}")
#pipe.load_lora_weights(lora_path, adapter_name = "eden_lora")#, weight_name="pytorch_lora_weights.safetensors")
#scales = {...}
#pipe.set_adapters("my_adapter", scales)
#pipe.set_adapters("eden_lora", scales)
for key in list_adapters_component_wise:
adapter_names = list_adapters_component_wise[key]
for adapter_name in adapter_names:
print(f"Set adapter '{adapter_name}' of '{key}' with scale = {lora_scale}")
pipe.set_adapters(adapter_name, adapter_weights=[lora_scale])
pipe.unet.merge_adapter()
# Load the textual_inversion token embeddings into the pipeline:
try: #SDXL
@@ -123,9 +112,10 @@ def save_lora(
elif len(unet_lora_parameters) > 0:
unet.save_pretrained(save_directory = output_dir)
if 0:
unet_lora_layers_to_save = convert_state_dict_to_diffusers(get_peft_model_state_dict(unet))
lora_tensors = get_peft_model_state_dict(unet)
if 1:
unet_lora_layers_to_save = convert_state_dict_to_diffusers(lora_tensors)
StableDiffusionXLPipeline.save_lora_weights(
output_dir,
unet_lora_layers=unet_lora_layers_to_save,
@@ -135,7 +125,6 @@ def save_lora(
if 1:
#lora_tensors = unet_attn_processors_state_dict(unet)
lora_tensors = get_peft_model_state_dict(unet)
save_file(lora_tensors, f"{output_dir}/{name}_lora_orig.safetensors")
embedding_handler.save_embeddings(f"{output_dir}/{name}_embeddings.safetensors")
-102
View File
@@ -1,102 +0,0 @@
import re
import json
import os
def replace_in_string(s, replacements):
while True:
replaced = False
for target, replacement in replacements.items():
new_s = re.sub(target, replacement, s, flags=re.IGNORECASE)
if new_s != s:
s = new_s
replaced = True
if not replaced:
break
return s
def prepare_prompt_for_lora(prompt, lora_path, trigger_text: str, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
orig_prompt = prompt
# Helper function to read JSON
def read_json_from_path(path):
with open(path, "r") as f:
return json.load(f)
# Check existence of "special_params.json"
if not os.path.exists(os.path.join(lora_path, "special_params.json")):
raise ValueError("This concept is from an old lora trainer that was deprecated. Please retrain your concept for better results!")
token_map = read_json_from_path(os.path.join(lora_path, "special_params.json"))
training_args = read_json_from_path(os.path.join(lora_path, "training_args.json"))
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
return trigger_text + ", " + prompt
lora_name_encapsulated = "<" + lora_name + ">"
trigger_text = trigger_text
try:
mode = training_args["concept_mode"]
except KeyError:
try:
mode = training_args["mode"]
except KeyError:
mode = "object"
# Handle different modes
if mode != "style":
replacements = {
"<concept>": trigger_text,
"<concepts>": trigger_text + "'s",
lora_name_encapsulated: trigger_text,
lora_name_encapsulated.lower(): trigger_text,
lora_name: trigger_text,
lora_name.lower(): trigger_text,
}
prompt = replace_in_string(prompt, replacements)
if trigger_text not in prompt:
prompt = trigger_text + ", " + prompt
else:
style_replacements = {
"in the style of <concept>": "in the style of TOK",
f"in the style of {lora_name_encapsulated}": "in the style of TOK",
f"in the style of {lora_name_encapsulated.lower()}": "in the style of TOK",
f"in the style of {lora_name}": "in the style of TOK",
f"in the style of {lora_name.lower()}": "in the style of TOK"
}
prompt = replace_in_string(prompt, style_replacements)
if "in the style of TOK" not in prompt:
prompt = "in the style of TOK, " + prompt
# Final cleanup
prompt = replace_in_string(prompt, {"<concept>": "TOK", lora_name_encapsulated: "TOK"})
if interpolation and mode != "style":
prompt = "TOK, " + prompt
# Replace tokens based on token map
prompt = replace_in_string(prompt, token_map)
# Fix common mistakes
fix_replacements = {
r",,": ",",
r"\s\s+": " ", # Replaces one or more whitespace characters with a single space
r"\s\.": ".",
r"\s,": ","
}
prompt = replace_in_string(prompt, fix_replacements)
if verbose:
print('-------------------------')
print("Adjusted prompt for LORA:")
print(orig_prompt)
print('-- to:')
print(prompt)
print('-------------------------')
return prompt
+3 -2
View File
@@ -250,7 +250,8 @@ def main(
safeguard_warmup=True,
weight_decay=config.lora_weight_decay if not config.use_dora else 0.0,
betas=(0.9, 0.99),
growth_rate=1.025, # this slows down the lr_rampup
#growth_rate=1.025, # this slows down the lr_rampup
growth_rate=1.05, # this slows down the lr_rampup
)
train_dataset = PreprocessedDataset(
@@ -560,7 +561,7 @@ def main(
ti_lrs.append(0.0)
# Print some statistics:
if config.debug and (global_step % config.checkpointing_steps == 0): # and global_step > 0:
if config.debug and (global_step % config.checkpointing_steps == 0) and global_step > 0:
output_save_dir = f"{checkpoint_dir}/checkpoint-{global_step}"
os.makedirs(output_save_dir, exist_ok=True)
config.save_as_json(