mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
Adding image generation container
This commit is contained in:
@@ -0,0 +1,266 @@
|
||||
import argparse
|
||||
import os
|
||||
import pathlib
|
||||
import torch
|
||||
|
||||
# Note please use type hints.
|
||||
from stablediffusionfe import StableDiffusionFE
|
||||
import random
|
||||
import string
|
||||
|
||||
def generate_random_string(length=6):
|
||||
# Choose from uppercase, lowercase letters and digits
|
||||
characters = string.ascii_letters + string.digits
|
||||
# Use random.choices to generate a list of random characters, then join them into a string
|
||||
random_string = ''.join(random.choices(characters, k=length))
|
||||
return random_string
|
||||
|
||||
class ModelInfo:
|
||||
def __init__(self,path:pathlib.Path):
|
||||
self.version = 1.5
|
||||
self.path = path
|
||||
|
||||
class StableDiffusion(ModelInfo):
|
||||
def __init__(self, path: pathlib.Path):
|
||||
super().__init__(path)
|
||||
|
||||
class loRA(ModelInfo):
|
||||
def __init__(self, path: pathlib.Path):
|
||||
super().__init__(path)
|
||||
|
||||
class VAE(ModelInfo):
|
||||
def __init__(self, path: pathlib.Path):
|
||||
super().__init__(path)
|
||||
|
||||
|
||||
_au111_to_diffusers_samplers = {'DPM++ 2M Karras': 'DPMSolverMultistepScheduler',
|
||||
'DPM++ 2M SDE Exponential': 'DPMSolverSDEScheduler',
|
||||
'DPM++ 2M SDE Karras': 'DPMSolverSDEScheduler',
|
||||
'Euler a': 'EulerAncestralDiscreteScheduler',
|
||||
'Euler': 'EulerDiscreteScheduler',
|
||||
'LMS': 'LMSDiscreteScheduler',
|
||||
'Heun': 'HeunDiscreteScheduler',
|
||||
'DPM2': 'KDPM2DiscreteScheduler',
|
||||
'DPM2 a': 'KDPM2AncestralDiscreteScheduler',
|
||||
'DPM++ 2S a': 'DDPMScheduler',
|
||||
'DPM++ 2M': 'DPMSolverMultistepScheduler',
|
||||
'DPM++ 2M SDE': 'DPMSolverSDEScheduler',
|
||||
'DPM++ 2M SDE Heun': 'DPMSolverSDEScheduler',
|
||||
'DPM++ 2M SDE Heun Karras': 'DPMSolverSDEScheduler',
|
||||
'DPM++ 2M SDE Heun Exponential': 'DPMSolverSDEScheduler',
|
||||
'DPM++ 3M SDE': 'UniPCMultistepScheduler',
|
||||
'DPM++ 3M SDE Karras': 'UniPCMultistepScheduler',
|
||||
'DPM++ 3M SDE Exponential': 'UniPCMultistepScheduler',
|
||||
'DPM fast': 'DEISMultistepScheduler',
|
||||
'DPM adaptive': 'DDIMScheduler',
|
||||
'LMS Karras': 'LMSDiscreteScheduler',
|
||||
'DPM2 Karras': 'KDPM2DiscreteScheduler',
|
||||
'DPM2 a Karras': 'KDPM2AncestralDiscreteScheduler',
|
||||
'DPM++ 2S a Karras': 'DDPMScheduler'}
|
||||
|
||||
|
||||
_available_samplers = ['DDIMScheduler', 'DPMSolverSDEScheduler', 'DDPMScheduler',
|
||||
'EulerDiscreteScheduler', 'DPMSolverMultistepScheduler',
|
||||
'DEISMultistepScheduler', 'LMSDiscreteScheduler', 'PNDMScheduler',
|
||||
'HeunDiscreteScheduler', 'KDPM2AncestralDiscreteScheduler', 'DPMSolverSinglestepScheduler',
|
||||
'KDPM2DiscreteScheduler', 'UniPCMultistepScheduler', 'EulerAncestralDiscreteScheduler']
|
||||
|
||||
class Sampler:
|
||||
samplers_k_diffusion_dict = {
|
||||
'DPM++ 2M Karras': ('DPM++ 2M Karras', 'sample_dpmpp_2m', ['k_dpmpp_2m_ka'], {'scheduler': 'karras'}),
|
||||
'DPM++ SDE Karras': ('DPM++ SDE Karras', 'sample_dpmpp_sde', ['k_dpmpp_sde_ka'], {'scheduler': 'karras', "second_order": True, "brownian_noise": True}),
|
||||
'DPM++ 2M SDE Exponential': ('DPM++ 2M SDE Exponential', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_exp'], {'scheduler': 'exponential', "brownian_noise": True}),
|
||||
'DPM++ 2M SDE Karras': ('DPM++ 2M SDE Karras', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_ka'], {'scheduler': 'karras', "brownian_noise": True}),
|
||||
'Euler a': ('Euler a', 'sample_euler_ancestral', ['k_euler_a', 'k_euler_ancestral'], {"uses_ensd": True}),
|
||||
'Euler': ('Euler', 'sample_euler', ['k_euler'], {}),
|
||||
'LMS': ('LMS', 'sample_lms', ['k_lms'], {}),
|
||||
'Heun': ('Heun', 'sample_heun', ['k_heun'], {"second_order": True}),
|
||||
'DPM2': ('DPM2', 'sample_dpm_2', ['k_dpm_2'], {'discard_next_to_last_sigma': True, "second_order": True}),
|
||||
'DPM2 a': ('DPM2 a', 'sample_dpm_2_ancestral', ['k_dpm_2_a'], {'discard_next_to_last_sigma': True, "uses_ensd": True, "second_order": True}),
|
||||
'DPM++ 2S a': ('DPM++ 2S a', 'sample_dpmpp_2s_ancestral', ['k_dpmpp_2s_a'], {"uses_ensd": True, "second_order": True}),
|
||||
'DPM++ 2M': ('DPM++ 2M', 'sample_dpmpp_2m', ['k_dpmpp_2m'], {}),
|
||||
'DPM++ SDE': ('DPM++ SDE', 'sample_dpmpp_sde', ['k_dpmpp_sde'], {"second_order": True, "brownian_noise": True}),
|
||||
'DPM++ 2M SDE': ('DPM++ 2M SDE', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_ka'], {"brownian_noise": True}),
|
||||
'DPM++ 2M SDE Heun': ('DPM++ 2M SDE Heun', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_heun'], {"brownian_noise": True, "solver_type": "heun"}),
|
||||
'DPM++ 2M SDE Heun Karras': ('DPM++ 2M SDE Heun Karras', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_heun_ka'], {'scheduler': 'karras', "brownian_noise": True, "solver_type": "heun"}),
|
||||
'DPM++ 2M SDE Heun Exponential': ('DPM++ 2M SDE Heun Exponential', 'sample_dpmpp_2m_sde', ['k_dpmpp_2m_sde_heun_exp'], {'scheduler': 'exponential', "brownian_noise": True, "solver_type": "heun"}),
|
||||
'DPM++ 3M SDE': ('DPM++ 3M SDE', 'sample_dpmpp_3m_sde', ['k_dpmpp_3m_sde'], {'discard_next_to_last_sigma': True, "brownian_noise": True}),
|
||||
'DPM++ 3M SDE Karras': ('DPM++ 3M SDE Karras', 'sample_dpmpp_3m_sde', ['k_dpmpp_3m_sde_ka'], {'scheduler': 'karras', 'discard_next_to_last_sigma': True, "brownian_noise": True}),
|
||||
'DPM++ 3M SDE Exponential': ('DPM++ 3M SDE Exponential', 'sample_dpmpp_3m_sde', ['k_dpmpp_3m_sde_exp'], {'scheduler': 'exponential', 'discard_next_to_last_sigma': True, "brownian_noise": True}),
|
||||
'DPM fast': ('DPM fast', 'sample_dpm_fast', ['k_dpm_fast'], {"uses_ensd": True}),
|
||||
'DPM adaptive': ('DPM adaptive', 'sample_dpm_adaptive', ['k_dpm_ad'], {"uses_ensd": True}),
|
||||
'LMS Karras': ('LMS Karras', 'sample_lms', ['k_lms_ka'], {'scheduler': 'karras'}),
|
||||
'DPM2 Karras': ('DPM2 Karras', 'sample_dpm_2', ['k_dpm_2_ka'], {'scheduler': 'karras', 'discard_next_to_last_sigma': True, "uses_ensd": True, "second_order": True}),
|
||||
'DPM2 a Karras': ('DPM2 a Karras', 'sample_dpm_2_ancestral', ['k_dpm_2_a_ka'], {'scheduler': 'karras', 'discard_next_to_last_sigma': True, "uses_ensd": True, "second_order": True}),
|
||||
'DPM++ 2S a Karras': ('DPM++ 2S a Karras', 'sample_dpmpp_2s_ancestral', ['k_dpmpp_2s_a_ka'], {'scheduler': 'karras', "uses_ensd": True, "second_order": True})
|
||||
}
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
class InputCheckers:
|
||||
@staticmethod
|
||||
def check_power_of_2(value:int):
|
||||
value = int(value)
|
||||
if value < 256 or value > 4096:
|
||||
raise argparse.ArgumentTypeError(f"Value {value} not in range [256, 4096]")
|
||||
if (value & (value - 1)) != 0:
|
||||
raise argparse.ArgumentTypeError(f"Value {value} is not a power of 2")
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def check_cfg_scale(value:int):
|
||||
value = int(value)
|
||||
if value < 1 or value > 30:
|
||||
raise argparse.ArgumentTypeError(f"Value {value} not in range [1, 30]")
|
||||
return value
|
||||
|
||||
@staticmethod
|
||||
def check_sampler(value:str, dictionary:[str,tuple]):
|
||||
return value in dictionary
|
||||
|
||||
@staticmethod
|
||||
def valid_path(path:pathlib.Path):
|
||||
if not os.path.exists(path):
|
||||
return path # TODO remove this later.
|
||||
raise argparse.ArgumentTypeError(f"The path '{path}' does not exist.")
|
||||
return path
|
||||
|
||||
def main():
|
||||
|
||||
parser = argparse.ArgumentParser(description="CLI for various inputs.")
|
||||
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)
|
||||
# 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("--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,
|
||||
help="Path to the OpenPose ControlNet")
|
||||
parser.add_argument("--openpose-aux", type=InputCheckers.valid_path, required=False,
|
||||
help="Path to the OpenPose auxiliary to the ControlNet (image->pose)")
|
||||
parser.add_argument("--lcm", type=InputCheckers.valid_path, required=False,
|
||||
help="Path to LCM LoRA")
|
||||
parser.add_argument("--image-input", type=InputCheckers.valid_path, required=False,
|
||||
help="Input image when using ControlNet")
|
||||
parser.add_argument("--batch-size",type=int,required=True, help="pick 1 by default")
|
||||
parser.add_argument("--batch-count", type=int, required=True, help="pick 1 by default")
|
||||
|
||||
parser.add_argument("--output-name", type=str, required=False, help="base name of output files. if not specified will be random string")
|
||||
parser.add_argument("--output-dir", type=str, required=False,
|
||||
help="output directory. default is outputs in working dir",
|
||||
default="outputs")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Logic here to process the arguments as needed
|
||||
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"width: {args.width}")
|
||||
print(f"height: {args.height}")
|
||||
print(f"cfg_scale: {args.cfg_scale}")
|
||||
print(f"seed: {args.seed}")
|
||||
print(f"loRA_path: {args.loRA_path}")
|
||||
print(f"check_point: {args.check_point}")
|
||||
print(f"vae: {args.vae}")
|
||||
print(f"batch_size: {args.batch_size}")
|
||||
print(f"batch_count: {args.batch_count}")
|
||||
|
||||
####### Load model ###########################
|
||||
# Set the device to CUDA if available, else CPU
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
# Initialize the StableDiffusionFE with the chosen device
|
||||
sd_fe = StableDiffusionFE(device=device)
|
||||
|
||||
# Load the model
|
||||
model_path = args.check_point
|
||||
sd_fe.load_model(model_path=model_path)
|
||||
|
||||
# Set CFG scale
|
||||
sd_fe.cfg_scale = args.cfg_scale
|
||||
|
||||
# Load LoRA
|
||||
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.seed is not None:
|
||||
sd_fe.manual_seed(args.seed)
|
||||
|
||||
# Load VAE if specified (not tested)
|
||||
if args.vae is not None:
|
||||
sd_fe.load_vae(vae_path=args.vae)
|
||||
|
||||
# Load LCM LoRA if specified
|
||||
if args.lcm is not None:
|
||||
sd_fe.load_lcm_lora(args.lcm)
|
||||
|
||||
# Load OpenPose controlnet if specified
|
||||
if args.openpose is not None:
|
||||
openpose_auxnet_name = args.openpose_aux
|
||||
openpose_controlnet_name = args.openpose
|
||||
|
||||
# Infuse the model with ControlNet abilities.
|
||||
sd_fe.make_controlnet(openpose_controlnet_name)
|
||||
# The make_controlnet function replaces the pipeline and reloads the last loaded LoRA
|
||||
|
||||
# Add the OpenPose auxiliary net that automatically converts poses in regular images to OpenPose things. For other
|
||||
# ControlNets, you'll have to manually instantiate its auxnet and use set_controlnet_aux
|
||||
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)
|
||||
|
||||
|
||||
####### Load model ###########################
|
||||
|
||||
### Generate images############
|
||||
|
||||
|
||||
generated_images = sd_fe.generate_count(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
num_inference_steps=args.number_of_samples,
|
||||
batch_size=args.batch_size,
|
||||
batch_count=args.batch_count,
|
||||
dimensions=dimensions,
|
||||
image_input=args.image_input,
|
||||
)
|
||||
|
||||
output_path = "outputs" if args.output_dir is None else args.output_dir
|
||||
print(f"Generated {len(generated_images)} images. Saving to {output_path}")
|
||||
#### Save ###########
|
||||
|
||||
if not os.path.isdir(output_path):
|
||||
os.mkdir(output_path)
|
||||
|
||||
filename_base = generate_random_string(6) if args.output_name is None else args.output_name
|
||||
for i, img in enumerate(generated_images):
|
||||
img_raw_filename = f"{filename_base}_{i}.png"
|
||||
img.save(
|
||||
os.path.join(output_path,
|
||||
img_raw_filename)
|
||||
)
|
||||
|
||||
print("Done!")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
certifi==2022.12.7
|
||||
charset-normalizer==2.1.1
|
||||
filelock==3.9.0
|
||||
fsspec==2023.10.0
|
||||
huggingface-hub==0.18.0
|
||||
idna==3.4
|
||||
Jinja2==3.1.2
|
||||
MarkupSafe==2.1.2
|
||||
mpmath==1.3.0
|
||||
networkx==3.0
|
||||
numpy==1.24.1
|
||||
packaging==23.2
|
||||
Pillow==9.3.0
|
||||
PyYAML==6.0.1
|
||||
regex==2023.10.3
|
||||
requests==2.28.1
|
||||
ruff==0.1.3
|
||||
safetensors==0.4.0
|
||||
sympy==1.12
|
||||
tokenizers==0.13.3
|
||||
torch==2.1.0+cu118
|
||||
torchaudio==2.1.0+cu118
|
||||
torchvision==0.16.0+cu118
|
||||
tqdm==4.66.1
|
||||
transformers==4.30.2
|
||||
triton==2.1.0
|
||||
typing_extensions==4.4.0
|
||||
urllib3==1.26.13
|
||||
xformers==0.0.22.post7+cu118
|
||||
diffusers
|
||||
@@ -0,0 +1,439 @@
|
||||
from diffusers import StableDiffusionPipeline, AutoencoderKL
|
||||
import torch
|
||||
from typing import List, Tuple, Optional, Union
|
||||
from PIL import Image
|
||||
from diffusers import StableDiffusionControlNetPipeline, ControlNetModel
|
||||
from controlnet_aux import OpenposeDetector
|
||||
from diffusers.utils import load_image
|
||||
from diffusers import LCMScheduler
|
||||
|
||||
|
||||
class StableDiffusionFE:
|
||||
"""
|
||||
Wrapper around diffusers StableDiffusionPipeline, as if that wasn't easy enough.
|
||||
Supports loading LoRAs, and count-batch generation like automatic111's
|
||||
"""
|
||||
|
||||
def __init__(self, device: str, model_path: str = None, controlnet_path: str = None,
|
||||
tensor_dtype: torch.dtype = torch.float32):
|
||||
"""
|
||||
Initialize the StableDiffusionFE with the specified device and optionally load a model.
|
||||
|
||||
Args:
|
||||
device (str): The device to run the model on (e.g., 'cuda', 'cpu').
|
||||
model_path (str, optional): Path to the .safetensors file to load the model from. Defaults to None.
|
||||
controlnet_path (str, optional): Path to ControlNet model. If not None and model_path is specified, will load
|
||||
the model in ControlNet mode.
|
||||
tensor_dtype (torch.dtype, optional): Only effective if model_path is passed, the type to load the model in.
|
||||
Defaults to torch.float32 (full precision)
|
||||
|
||||
|
||||
"""
|
||||
self.use_auxnet = False
|
||||
self.device = device
|
||||
self.pipeline = None
|
||||
self.cfg_scale = 7.5
|
||||
self.generator = None
|
||||
self.controlnet_aux = None
|
||||
self.model_path = model_path
|
||||
self.tensor_dtype = tensor_dtype
|
||||
self.lora_path = None
|
||||
self.lcm_lora_path = None
|
||||
|
||||
if model_path is not None:
|
||||
if controlnet_path is not None:
|
||||
self.make_controlnet(controlnet_path, model_path)
|
||||
else:
|
||||
self.load_model(model_path, tensor_dtype)
|
||||
|
||||
def load_model(self, model_path: str, tensor_dtype: torch.dtype = torch.float32) -> None:
|
||||
"""
|
||||
Load the Stable Diffusion model from a specified safetensors file path.
|
||||
|
||||
Args:
|
||||
model_path (str): Path to the .safetensors file to load the model from.
|
||||
tensor_dtype (torch.dtype, optional): Only effective if model_path is passed, the type to load the model in.
|
||||
Defaults to torch.float32 (full precision)
|
||||
|
||||
"""
|
||||
# Loading the pipeline from a .safetensors file without a safety checker.
|
||||
self.pipeline = StableDiffusionPipeline.from_single_file(model_path, safety_checker=None,
|
||||
torch_dtype=tensor_dtype).to(self.device)
|
||||
# Turn off safety checker because I don't want black images. (safety_checker=None above sometimes isn't enough)
|
||||
self.pipeline.requires_safety_checker = False
|
||||
self.pipeline.safety_checker = None
|
||||
self.model_path = model_path
|
||||
self.tensor_dtype = tensor_dtype
|
||||
|
||||
def manual_seed(self, seed: int) -> None:
|
||||
"""
|
||||
Creates a generator with a fixed seed and stores it to be used for later inferences
|
||||
|
||||
Args:
|
||||
seed (int): The seed value to use for random number generation.
|
||||
"""
|
||||
self.generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
|
||||
def set_controlnet_aux(self, auxnet_class: object, use_auxnet: bool) -> None:
|
||||
"""
|
||||
Sets the auxiliary network (e.g., OpenPose detector) for the class.
|
||||
|
||||
Args:
|
||||
auxnet_class: The class object to be set as the auxiliary network.
|
||||
use_auxnet: Whether to use the aux net automatically. If not, just stores it as self.controlnet_aux
|
||||
"""
|
||||
self.use_auxnet = use_auxnet
|
||||
self.controlnet_aux = auxnet_class
|
||||
|
||||
def make_openpose(self, controlnet_model_name: str):
|
||||
"""
|
||||
Creates an OpenPose auxiliary controlnet image pre-processor.
|
||||
|
||||
The aux module is used to convert regular input pictures to OpenPose things automatically
|
||||
|
||||
Args:
|
||||
controlnet_model_name (str): The name of the OpenPose model
|
||||
"""
|
||||
openpose = OpenposeDetector.from_pretrained(controlnet_model_name).to(self.device)
|
||||
self.set_controlnet_aux(openpose, True)
|
||||
|
||||
def make_controlnet(self, controlnet_model_name: str, model_path: str = None,
|
||||
auto_reload_lora: bool = True) -> None:
|
||||
"""
|
||||
Replaces the pipeline by loading the model with a specified ControlNet.
|
||||
|
||||
Args:
|
||||
controlnet_model_name (str): The name of the ControlNet model to load.
|
||||
model_path (str, optional): Path to the .safetensors file to load the main Stable Diffusion model from.
|
||||
Defaults to using self.model_path (reloading the model) if None.
|
||||
auto_reload_lora (bool, optional): Automatically reload (last loaded) LoRA weights, since we are reloading the model
|
||||
Defaults to True, but only effective if model_path is None
|
||||
|
||||
Note: If you're using OpenPose you can instantiate its aux-net to be used automatically with make_openpose
|
||||
"""
|
||||
reloading = False
|
||||
if model_path is None:
|
||||
# If no model was loaded...
|
||||
if self.model_path is None:
|
||||
raise ValueError("Must have a model loaded to use make_controlnet without a new model path")
|
||||
|
||||
reloading = True
|
||||
model_path = self.model_path
|
||||
|
||||
controlnet = ControlNetModel.from_pretrained(
|
||||
controlnet_model_name,
|
||||
torch_dtype=self.tensor_dtype
|
||||
).to(self.device)
|
||||
|
||||
self.pipeline = StableDiffusionControlNetPipeline.from_single_file(
|
||||
model_path,
|
||||
controlnet=controlnet,
|
||||
safety_checker=None,
|
||||
torch_dtype=self.tensor_dtype
|
||||
).to(self.device)
|
||||
|
||||
# Turn off safety checker because I don't want black images. (safety_checker=None above sometimes isn't enough)
|
||||
self.pipeline.requires_safety_checker = False
|
||||
self.pipeline.safety_checker = None
|
||||
self.model_path = model_path
|
||||
|
||||
if self.lora_path is not None and auto_reload_lora and reloading:
|
||||
self.load_lora_weights(self.lora_path)
|
||||
|
||||
if reloading and self.lcm_lora_path is not None:
|
||||
self.load_lcm_lora(self.lcm_lora_path)
|
||||
|
||||
|
||||
def load_lora_weights(self, lora_path: str) -> None:
|
||||
"""
|
||||
Load LoRA weights from a specified safetensors file path.
|
||||
|
||||
Args:
|
||||
lora_path (str): Path to the .safetensors file to load the LoRA weights from.
|
||||
"""
|
||||
if not self.pipeline:
|
||||
raise ValueError("Model must be loaded before loading LoRA weights.")
|
||||
|
||||
self.lora_path = lora_path
|
||||
self.pipeline.load_lora_weights(lora_path)
|
||||
|
||||
def load_vae(self, vae_path: str) -> None:
|
||||
"""
|
||||
Load a VAE model from the specified path and replace the existing VAE in the pipeline.
|
||||
|
||||
Args:
|
||||
vae_path (str): The file path to the VAE model.
|
||||
"""
|
||||
if not self.pipeline:
|
||||
raise ValueError("Pipeline must be initialized before loading a VAE.")
|
||||
|
||||
vae_model = AutoencoderKL.from_single_file(vae_path)
|
||||
self.pipeline.vae = vae_model.to(self.device) # Move the VAE to the correct device
|
||||
|
||||
def generate_image(
|
||||
self,
|
||||
prompt: str,
|
||||
num_inference_steps: int,
|
||||
negative_prompt: str = None,
|
||||
dimensions: Optional[Tuple[int, int]] = None,
|
||||
image_input: Optional[Union[str, Image.Image]] = None
|
||||
) -> Image.Image:
|
||||
"""
|
||||
Generate an image based on the prompt and number of inference steps provided.
|
||||
Optionally uses an image input for ControlNet.
|
||||
|
||||
Args:
|
||||
prompt (str): The prompt describing the desired output image.
|
||||
num_inference_steps (int): Number of inference steps to run the generation.
|
||||
negative_prompt (str, optional): Negative prompt.
|
||||
dimensions (tuple (int, int), optional): Width and height of output image for generation.
|
||||
image_input (Union[str, Image.Image], optional): The input image for ControlNet,
|
||||
can be a file path or a PIL.Image object.
|
||||
|
||||
Returns:
|
||||
Image.Image: The generated image.
|
||||
"""
|
||||
|
||||
if not self.pipeline:
|
||||
raise ValueError("Model must be loaded before generating images.")
|
||||
|
||||
width, height = dimensions if dimensions is not None else (None, None)
|
||||
|
||||
# Process the image input if provided
|
||||
if image_input is not None:
|
||||
image_input = self.process_image_input(image_input)
|
||||
|
||||
return self.pipeline(
|
||||
prompt,
|
||||
image_input,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=self.cfg_scale,
|
||||
generator=self.generator,
|
||||
negative_prompt=negative_prompt,
|
||||
width=width,
|
||||
height=height
|
||||
).images[0]
|
||||
else:
|
||||
# Normal image generation without ControlNet
|
||||
return self.pipeline(
|
||||
prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=self.cfg_scale,
|
||||
generator=self.generator,
|
||||
width=width,
|
||||
height=height
|
||||
).images[0]
|
||||
|
||||
def generate_images(
|
||||
self,
|
||||
prompts: List[str],
|
||||
negative_prompts: List[str],
|
||||
num_inference_steps: int,
|
||||
dimensions: Optional[Tuple[int, int]] = None,
|
||||
image_inputs: Optional[List[Union[str, Image.Image]]] = None,
|
||||
force_no_process_img: bool = False
|
||||
) -> List:
|
||||
"""
|
||||
Generate a batch of images based on the given lists of prompts and negative prompts,
|
||||
with the specified number of inference steps for each generation.
|
||||
Optionally uses image inputs for ControlNet.
|
||||
|
||||
Args:
|
||||
prompts (List[str]): The prompts describing the desired output images.
|
||||
negative_prompts (List[str]): The negative prompts to guide the generation away from undesired elements.
|
||||
num_inference_steps (int): Number of inference steps to run the generation.
|
||||
dimensions (tuple (int, int), optional): Width and height of output image for generation.
|
||||
image_inputs (List[Union[str, Image.Image]], optional): List of input images for ControlNet,
|
||||
can be file paths or PIL.Image objects.
|
||||
force_no_process_img (bool): If True, skips processing of image inputs with aux ControlNet
|
||||
|
||||
Returns:
|
||||
List: A list of generated images.
|
||||
"""
|
||||
|
||||
if not self.pipeline:
|
||||
raise ValueError("Model must be loaded before generating images.")
|
||||
if len(prompts) != len(negative_prompts):
|
||||
raise ValueError("The number of prompts and negative prompts must be the same.")
|
||||
if image_inputs is not None and len(prompts) != len(image_inputs):
|
||||
raise ValueError("The number of prompts and image inputs must be the same.")
|
||||
|
||||
width, height = dimensions if dimensions is not None else (None, None)
|
||||
|
||||
# Process image inputs if provided and self.controlnet_aux is not None
|
||||
processed_images = None
|
||||
if image_inputs is not None:
|
||||
processed_images = []
|
||||
|
||||
use_au_net = self.use_auxnet and self.controlnet_aux is not None
|
||||
use_controlnet_aux = use_au_net and not force_no_process_img
|
||||
|
||||
for img_input in image_inputs:
|
||||
if isinstance(img_input, str):
|
||||
img_input = load_image(img_input)
|
||||
|
||||
processed_image = self.controlnet_aux(img_input) if use_controlnet_aux else img_input
|
||||
processed_images.append(processed_image)
|
||||
|
||||
# Dynamic pipeline call based on whether image inputs are provided
|
||||
if processed_images is not None:
|
||||
images, _ = self.pipeline(
|
||||
prompts,
|
||||
image_input=processed_images,
|
||||
negative_prompt=negative_prompts,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=self.cfg_scale,
|
||||
generator=self.generator,
|
||||
width=width,
|
||||
height=height,
|
||||
return_dict=False
|
||||
)
|
||||
else:
|
||||
images, _ = self.pipeline(
|
||||
prompts,
|
||||
negative_prompt=negative_prompts,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=self.cfg_scale,
|
||||
generator=self.generator,
|
||||
width=width,
|
||||
height=height,
|
||||
return_dict=False
|
||||
)
|
||||
|
||||
return images
|
||||
|
||||
# Additional helper method for processing image inputs
|
||||
def process_image_input(self, img_input: Union[str, Image.Image]) -> Image.Image:
|
||||
if isinstance(img_input, str):
|
||||
img_input = load_image(img_input)
|
||||
|
||||
# Just to be sure. It should be impossible for self.use_auxnet to be True while self.controlnet_aux is None,
|
||||
# but eh.
|
||||
use_au_net = self.use_auxnet and self.controlnet_aux is not None
|
||||
return self.controlnet_aux(img_input) if use_au_net else img_input
|
||||
|
||||
def generate_count(
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
num_inference_steps: int,
|
||||
batch_size: int,
|
||||
batch_count: int,
|
||||
dimensions: Optional[Tuple[int, int]] = None,
|
||||
image_input: Optional[Union[str, Image.Image]] = None
|
||||
|
||||
) -> List:
|
||||
"""
|
||||
Generate a specific number of images based on a single prompt and negative prompt,
|
||||
processing in batches of a given size.
|
||||
|
||||
Args:
|
||||
prompt (str): The prompt describing the desired output image.
|
||||
negative_prompt (str): The negative prompt to guide the generation away from undesired elements.
|
||||
num_inference_steps (int): Number of inference steps to run the generation.
|
||||
batch_size (int): Number of images to generate in one batch.
|
||||
batch_count (int): Total number of images to generate.
|
||||
dimensions (tuple (int, int)): width and height of output image for generation
|
||||
image_input (Union[str, Image.Image], optional): A single input image for ControlNet,
|
||||
can be a file path or a PIL.Image object.
|
||||
|
||||
Returns:
|
||||
List: A list of generated images.
|
||||
"""
|
||||
|
||||
if not self.pipeline:
|
||||
raise ValueError("Model must be loaded before generating images.")
|
||||
|
||||
# Prepare lists to hold the full batch count of prompts and negative prompts
|
||||
full_prompts = [prompt] * batch_count
|
||||
full_negative_prompts = [negative_prompt] * batch_count
|
||||
|
||||
# Process the single image input once if provided
|
||||
processed_image = self.process_image_input(image_input) if image_input is not None else None
|
||||
|
||||
all_images = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, batch_count, batch_size):
|
||||
# Determine the size of the current batch
|
||||
current_batch_size = min(batch_size, batch_count - i)
|
||||
batch_prompts = full_prompts[i:i + current_batch_size]
|
||||
batch_negative_prompts = full_negative_prompts[i:i + current_batch_size]
|
||||
|
||||
# Generate images for the current batch
|
||||
batch_images = self.generate_images(
|
||||
prompts=batch_prompts,
|
||||
negative_prompts=batch_negative_prompts,
|
||||
num_inference_steps=num_inference_steps,
|
||||
dimensions=dimensions,
|
||||
image_inputs=[processed_image] * current_batch_size if processed_image is not None else None,
|
||||
force_no_process_img=True,
|
||||
)
|
||||
all_images.extend(batch_images)
|
||||
|
||||
return all_images
|
||||
|
||||
def get_compatible_schedulers(self) -> List[str]:
|
||||
"""
|
||||
Returns a list of compatible scheduler class names as strings,
|
||||
formatted without the 'diffusers.schedulers.scheduling_' prefix.
|
||||
|
||||
Returns:
|
||||
List[str]: A list of scheduler names.
|
||||
"""
|
||||
if not self.pipeline:
|
||||
raise ValueError("Pipeline must be loaded before getting schedulers.")
|
||||
|
||||
# Extract class names and format them
|
||||
scheduler_names = [
|
||||
cls.__name__.replace('Scheduling', '').replace('_', '')
|
||||
for cls in self.pipeline.scheduler.compatibles
|
||||
]
|
||||
return scheduler_names
|
||||
|
||||
def replace_scheduler(self, scheduler_name: str) -> None:
|
||||
"""
|
||||
Instantiates and replaces the scheduler in the pipeline with the given scheduler name.
|
||||
|
||||
Args:
|
||||
scheduler_name (str): The name of the scheduler to instantiate.
|
||||
"""
|
||||
if not self.pipeline:
|
||||
raise ValueError("Pipeline must be loaded before replacing scheduler.")
|
||||
|
||||
# Mapping of scheduler name to the class in diffusers
|
||||
schedulers_mapping = {
|
||||
cls.__name__.replace('Scheduling', '').replace('_', ''): cls
|
||||
for cls in self.pipeline.scheduler.compatibles
|
||||
}
|
||||
|
||||
# Find the scheduler class from the given scheduler name
|
||||
scheduler_class = schedulers_mapping.get(scheduler_name)
|
||||
if not scheduler_class:
|
||||
raise ValueError(f"Scheduler '{scheduler_name}' is not recognized as a compatible scheduler.")
|
||||
|
||||
# Instantiate and replace the scheduler in the pipeline
|
||||
self.pipeline.scheduler = scheduler_class.from_config(self.pipeline.scheduler.config)
|
||||
|
||||
def load_lcm_lora(self, adapter_id: str) -> None:
|
||||
"""
|
||||
Loads a Latent Consistency Model (LCM) with LoRA, allowing generation in fewer steps.
|
||||
|
||||
Args:
|
||||
adapter_id (str): The identifier of the LCM LoRA adapter.
|
||||
"""
|
||||
if not self.pipeline:
|
||||
raise ValueError("Pipeline must be loaded before loading LCM LoRA.")
|
||||
|
||||
# Update the scheduler to LCMScheduler
|
||||
self.pipeline.scheduler = LCMScheduler.from_config(self.pipeline.scheduler.config)
|
||||
|
||||
# Load and fuse LCM LoRA
|
||||
self.pipeline.load_lora_weights(adapter_id)
|
||||
self.pipeline.fuse_lora()
|
||||
|
||||
# https://huggingface.co/latent-consistency/lcm-lora-sdv1-5
|
||||
# >Please make sure to either disable guidance_scale or use values between 1.0 and 2.0.
|
||||
self.cfg_scale = 0
|
||||
self.lcm_lora_path = adapter_id
|
||||
+14
@@ -0,0 +1,14 @@
|
||||
python3.10 main.py \
|
||||
--prompt "This is a prompt" \
|
||||
--negative-prompt "This is a negative prompt" \
|
||||
--number-of-samples 20 \
|
||||
--samplers "DPM++ 2M Karras" \
|
||||
--width 512 \
|
||||
--height 256 \
|
||||
--cfg-scale 10 \
|
||||
--seed 1234 \
|
||||
--loRA-path "/path/to/loRA" \
|
||||
--check-point "/path/to/checkpoint" \
|
||||
--vae "/path/to/vae" \
|
||||
--batch-size 1 \
|
||||
--batch-count 1
|
||||
Reference in New Issue
Block a user