From 7d8b7b2c63ab879dd4e6a18e2a5780cb52e3018e Mon Sep 17 00:00:00 2001 From: mayukhdeb Date: Fri, 29 Mar 2024 12:36:28 -0700 Subject: [PATCH] pass trigger_text as arg --- trainer/utils/inference.py | 4 ++-- trainer/utils/prompt.py | 6 +++--- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/trainer/utils/inference.py b/trainer/utils/inference.py index 4cfda4c..7968ff9 100644 --- a/trainer/utils/inference.py +++ b/trainer/utils/inference.py @@ -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", diff --git a/trainer/utils/prompt.py b/trainer/utils/prompt.py index acd8086..2a4a62b 100644 --- a/trainer/utils/prompt.py +++ b/trainer/utils/prompt.py @@ -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"]