Merge pull request #28 from storytold/feature/image_generation

Adding image generation container
This commit is contained in:
ArEnSc
2023-12-30 12:44:13 -05:00
committed by GitHub
4 changed files with 749 additions and 0 deletions
@@ -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
View File
@@ -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