From 2bc24a438404aceadd719b6a0985562cd80d5bdc Mon Sep 17 00:00:00 2001 From: ArEnSc Date: Sat, 30 Dec 2023 12:31:12 -0500 Subject: [PATCH] Adding image generation container --- .../Storyteller-Studio-Filter/main.py | 266 +++++++++++ .../requirements.txt | 30 ++ .../stablediffusionfe.py | 439 ++++++++++++++++++ .../Storyteller-Studio-Filter/test.sh | 14 + 4 files changed, 749 insertions(+) create mode 100644 image_generation/Storyteller-Studio-Filter/main.py create mode 100644 image_generation/Storyteller-Studio-Filter/requirements.txt create mode 100644 image_generation/Storyteller-Studio-Filter/stablediffusionfe.py create mode 100755 image_generation/Storyteller-Studio-Filter/test.sh diff --git a/image_generation/Storyteller-Studio-Filter/main.py b/image_generation/Storyteller-Studio-Filter/main.py new file mode 100644 index 0000000..218df91 --- /dev/null +++ b/image_generation/Storyteller-Studio-Filter/main.py @@ -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() + + diff --git a/image_generation/Storyteller-Studio-Filter/requirements.txt b/image_generation/Storyteller-Studio-Filter/requirements.txt new file mode 100644 index 0000000..d8e0755 --- /dev/null +++ b/image_generation/Storyteller-Studio-Filter/requirements.txt @@ -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 diff --git a/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py b/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py new file mode 100644 index 0000000..3420946 --- /dev/null +++ b/image_generation/Storyteller-Studio-Filter/stablediffusionfe.py @@ -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 \ No newline at end of file diff --git a/image_generation/Storyteller-Studio-Filter/test.sh b/image_generation/Storyteller-Studio-Filter/test.sh new file mode 100755 index 0000000..0268912 --- /dev/null +++ b/image_generation/Storyteller-Studio-Filter/test.sh @@ -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