mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
Fixing The Inference Script for SD
This commit is contained in:
@@ -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:
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user