pass trigger_text as arg

This commit is contained in:
mayukhdeb
2024-03-29 12:36:28 -07:00
parent d2b30450e4
commit 7d8b7b2c63
2 changed files with 5 additions and 5 deletions
+2 -2
View File
@@ -11,7 +11,7 @@ from .prompt import prepare_prompt_for_lora
from .io import make_validation_img_grid
@torch.no_grad()
def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0", verbose: bool = True):
def render_images(training_pipeline, 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):
random.seed(seed)
@@ -55,7 +55,7 @@ def render_images(training_pipeline, render_size, lora_path, train_step, seed, i
training_scheduler = pipeline.scheduler
pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config)
validation_prompts = [prepare_prompt_for_lora(prompt, lora_path, verbose=verbose) for prompt in validation_prompts_raw]
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)
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",
+3 -3
View File
@@ -14,7 +14,7 @@ def replace_in_string(s, replacements):
break
return s
def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True):
def prepare_prompt_for_lora(prompt, lora_path, trigger_text: str, interpolation=False, verbose=True):
if "_no_token" in lora_path:
return prompt
@@ -35,10 +35,10 @@ def prepare_prompt_for_lora(prompt, lora_path, interpolation=False, verbose=True
try:
lora_name = str(training_args["name"])
except: # fallback for old loras that dont have the name field:
return training_args["trigger_text"] + ", " + prompt
return trigger_text + ", " + prompt
lora_name_encapsulated = "<" + lora_name + ">"
trigger_text = training_args["trigger_text"]
trigger_text = trigger_text
try:
mode = training_args["concept_mode"]