inference updates
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user