Fixing The Inference Script for SD

This commit is contained in:
ArEnSc
2024-01-14 17:15:53 -05:00
parent 2bc24a4384
commit 2efc940bf0
2 changed files with 14 additions and 12 deletions
@@ -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)
@@ -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:
"""