From 2efc940bf0cea0d5d122bb0b2c63e24936d2f625 Mon Sep 17 00:00:00 2001 From: ArEnSc Date: Sun, 14 Jan 2024 17:15:53 -0500 Subject: [PATCH] Fixing The Inference Script for SD --- .../Storyteller-Studio-Filter/main.py | 23 ++++++++++--------- .../stablediffusionfe.py | 3 ++- 2 files changed, 14 insertions(+), 12 deletions(-) diff --git a/image_generation/Storyteller-Studio-Filter/main.py b/image_generation/Storyteller-Studio-Filter/main.py index 218df91..930dced 100644 --- a/image_generation/Storyteller-Studio-Filter/main.py +++ b/image_generation/Storyteller-Studio-Filter/main.py @@ -131,14 +131,16 @@ def main(): parser.add_argument("--prompt", type=str, required=True, help="Text input called prompt") parser.add_argument("--negative-prompt", type=str, required=True, help="Text input called negative prompt") parser.add_argument("--number-of-samples",type=int, required=True) - parser.add_argument("--sampler", type=str, required=False, help="Text input called samplers", choices=_available_samplers) + #parser.add_argument("--sampler", type=str, required=False, help="Text input called samplers", choices=_available_samplers) + parser.add_argument("--samplers", type=str, required=False, help="Text input called samplers", choices=list(_au111_to_diffusers_samplers.keys())) + # ZDisket: We shoud really keep these static and instead use an upsampler afterwards instead parser.add_argument("--width", type=InputCheckers.check_power_of_2, required=False, help="Width value (power of 2, between 256 and 4096)") parser.add_argument("--height", type=InputCheckers.check_power_of_2, required=False, help="Height value (power of 2, between 256 and 4096)") - + parser.add_argument("--cfg-scale", type=InputCheckers.check_cfg_scale, required=True, help="Integer input called CFG scale (range from 1-30)") parser.add_argument("--seed", type=int, required=False, help="Seed value (integer of any value)") - parser.add_argument("--loRA-path", type=InputCheckers.valid_path, required=True, help="Path object to the loRA") + parser.add_argument("--loRA-path", type=InputCheckers.valid_path, required=False, help="Path object to the loRA") parser.add_argument("--check-point", type=InputCheckers.valid_path, required=True, help="Path object to the Check Point") parser.add_argument("--vae", type=InputCheckers.valid_path, required=False, help="Path to the Variational auto encoder") parser.add_argument("--openpose", type=InputCheckers.valid_path, required=False, @@ -163,7 +165,7 @@ def main(): print(f"prompt: {args.prompt}") print(f"negative_prompt: {args.negative_prompt}") print(f"number of samples: {args.number_of_samples}") - print(f"samplers: {args.sampler}") + print(f"samplers: {args.samplers}") print(f"width: {args.width}") print(f"height: {args.height}") print(f"cfg_scale: {args.cfg_scale}") @@ -189,13 +191,14 @@ def main(): sd_fe.cfg_scale = args.cfg_scale # Load LoRA - lora_path = args.loRA_path - sd_fe.load_lora_weights(lora_path=lora_path) + if args.loRA_path is not None: + lora_path = args.loRA_path + sd_fe.load_lora_weights(lora_path=lora_path) # Replace sampler if specified - if args.sampler is not None: - sd_fe.replace_scheduler(args.sampler) - + if args.samplers is not None: + sd_fe.replace_scheduler(_au111_to_diffusers_samplers[args.samplers]) + if args.seed is not None: sd_fe.manual_seed(args.seed) @@ -221,8 +224,6 @@ def main(): if openpose_auxnet_name is not None: sd_fe.make_openpose(openpose_auxnet_name) - - # Set width and height to dimensions. If any or both are None, it will result in the pipeline doing its default. dimensions = (args.width, args.height) diff --git a/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py b/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py index 3420946..d87a610 100644 --- a/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py +++ b/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py @@ -155,7 +155,8 @@ class StableDiffusionFE: raise ValueError("Model must be loaded before loading LoRA weights.") self.lora_path = lora_path - self.pipeline.load_lora_weights(lora_path) + + self.pipeline.load_lora_weights(lora_path,low_cpu_mem_usage=False, ignore_mismatched_sizes=True) def load_vae(self, vae_path: str) -> None: """