zparadox/stable-video-diffusion
0
1import pathlib2from dataclasses import asdict, dataclass3from enum import Enum4from typing import Optional5 6from omegaconf import OmegaConf7 8from sgm.inference.helpers import (Img2ImgDiscretizationWrapper, do_img2img,9 do_sample)10from sgm.modules.diffusionmodules.sampling import (DPMPP2MSampler,11 DPMPP2SAncestralSampler,12 EulerAncestralSampler,13 EulerEDMSampler,14 HeunEDMSampler,15 LinearMultistepSampler)16from sgm.util import load_model_from_config17 18 19class ModelArchitecture(str, Enum):20 SD_2_1 = "stable-diffusion-v2-1"21 SD_2_1_768 = "stable-diffusion-v2-1-768"22 SDXL_V0_9_BASE = "stable-diffusion-xl-v0-9-base"23 SDXL_V0_9_REFINER = "stable-diffusion-xl-v0-9-refiner"24 SDXL_V1_BASE = "stable-diffusion-xl-v1-base"25 SDXL_V1_REFINER = "stable-diffusion-xl-v1-refiner"26 27 28class Sampler(str, Enum):29 EULER_EDM = "EulerEDMSampler"30 HEUN_EDM = "HeunEDMSampler"31 EULER_ANCESTRAL = "EulerAncestralSampler"32 DPMPP2S_ANCESTRAL = "DPMPP2SAncestralSampler"33 DPMPP2M = "DPMPP2MSampler"34 LINEAR_MULTISTEP = "LinearMultistepSampler"35 36 37class Discretization(str, Enum):38 LEGACY_DDPM = "LegacyDDPMDiscretization"39 EDM = "EDMDiscretization"40 41 42class Guider(str, Enum):43 VANILLA = "VanillaCFG"44 IDENTITY = "IdentityGuider"45 46 47class Thresholder(str, Enum):48 NONE = "None"49 50 51@dataclass52class SamplingParams:53 width: int = 102454 height: int = 102455 steps: int = 5056 sampler: Sampler = Sampler.DPMPP2M57 discretization: Discretization = Discretization.LEGACY_DDPM58 guider: Guider = Guider.VANILLA59 thresholder: Thresholder = Thresholder.NONE60 scale: float = 6.061 aesthetic_score: float = 5.062 negative_aesthetic_score: float = 5.063 img2img_strength: float = 1.064 orig_width: int = 102465 orig_height: int = 102466 crop_coords_top: int = 067 crop_coords_left: int = 068 sigma_min: float = 0.029269 sigma_max: float = 14.614670 rho: float = 3.071 s_churn: float = 0.072 s_tmin: float = 0.073 s_tmax: float = 999.074 s_noise: float = 1.075 eta: float = 1.076 order: int = 477 78 79@dataclass80class SamplingSpec:81 width: int82 height: int83 channels: int84 factor: int85 is_legacy: bool86 config: str87 ckpt: str88 is_guided: bool89 90 91model_specs = {92 ModelArchitecture.SD_2_1: SamplingSpec(93 height=512,94 width=512,95 channels=4,96 factor=8,97 is_legacy=True,98 config="sd_2_1.yaml",99 ckpt="v2-1_512-ema-pruned.safetensors",100 is_guided=True,101 ),102 ModelArchitecture.SD_2_1_768: SamplingSpec(103 height=768,104 width=768,105 channels=4,106 factor=8,107 is_legacy=True,108 config="sd_2_1_768.yaml",109 ckpt="v2-1_768-ema-pruned.safetensors",110 is_guided=True,111 ),112 ModelArchitecture.SDXL_V0_9_BASE: SamplingSpec(113 height=1024,114 width=1024,115 channels=4,116 factor=8,117 is_legacy=False,118 config="sd_xl_base.yaml",119 ckpt="sd_xl_base_0.9.safetensors",120 is_guided=True,121 ),122 ModelArchitecture.SDXL_V0_9_REFINER: SamplingSpec(123 height=1024,124 width=1024,125 channels=4,126 factor=8,127 is_legacy=True,128 config="sd_xl_refiner.yaml",129 ckpt="sd_xl_refiner_0.9.safetensors",130 is_guided=True,131 ),132 ModelArchitecture.SDXL_V1_BASE: SamplingSpec(133 height=1024,134 width=1024,135 channels=4,136 factor=8,137 is_legacy=False,138 config="sd_xl_base.yaml",139 ckpt="sd_xl_base_1.0.safetensors",140 is_guided=True,141 ),142 ModelArchitecture.SDXL_V1_REFINER: SamplingSpec(143 height=1024,144 width=1024,145 channels=4,146 factor=8,147 is_legacy=True,148 config="sd_xl_refiner.yaml",149 ckpt="sd_xl_refiner_1.0.safetensors",150 is_guided=True,151 ),152}153 154 155class SamplingPipeline:156 def __init__(157 self,158 model_id: ModelArchitecture,159 model_path="checkpoints",160 config_path="configs/inference",161 device="cuda",162 use_fp16=True,163 ) -> None:164 if model_id not in model_specs:165 raise ValueError(f"Model {model_id} not supported")166 self.model_id = model_id167 self.specs = model_specs[self.model_id]168 self.config = str(pathlib.Path(config_path, self.specs.config))169 self.ckpt = str(pathlib.Path(model_path, self.specs.ckpt))170 self.device = device171 self.model = self._load_model(device=device, use_fp16=use_fp16)172 173 def _load_model(self, device="cuda", use_fp16=True):174 config = OmegaConf.load(self.config)175 model = load_model_from_config(config, self.ckpt)176 if model is None:177 raise ValueError(f"Model {self.model_id} could not be loaded")178 model.to(device)179 if use_fp16:180 model.conditioner.half()181 model.model.half()182 return model183 184 def text_to_image(185 self,186 params: SamplingParams,187 prompt: str,188 negative_prompt: str = "",189 samples: int = 1,190 return_latents: bool = False,191 ):192 sampler = get_sampler_config(params)193 value_dict = asdict(params)194 value_dict["prompt"] = prompt195 value_dict["negative_prompt"] = negative_prompt196 value_dict["target_width"] = params.width197 value_dict["target_height"] = params.height198 return do_sample(199 self.model,200 sampler,201 value_dict,202 samples,203 params.height,204 params.width,205 self.specs.channels,206 self.specs.factor,207 force_uc_zero_embeddings=["txt"] if not self.specs.is_legacy else [],208 return_latents=return_latents,209 filter=None,210 )211 212 def image_to_image(213 self,214 params: SamplingParams,215 image,216 prompt: str,217 negative_prompt: str = "",218 samples: int = 1,219 return_latents: bool = False,220 ):221 sampler = get_sampler_config(params)222 223 if params.img2img_strength < 1.0:224 sampler.discretization = Img2ImgDiscretizationWrapper(225 sampler.discretization,226 strength=params.img2img_strength,227 )228 height, width = image.shape[2], image.shape[3]229 value_dict = asdict(params)230 value_dict["prompt"] = prompt231 value_dict["negative_prompt"] = negative_prompt232 value_dict["target_width"] = width233 value_dict["target_height"] = height234 return do_img2img(235 image,236 self.model,237 sampler,238 value_dict,239 samples,240 force_uc_zero_embeddings=["txt"] if not self.specs.is_legacy else [],241 return_latents=return_latents,242 filter=None,243 )244 245 def refiner(246 self,247 params: SamplingParams,248 image,249 prompt: str,250 negative_prompt: Optional[str] = None,251 samples: int = 1,252 return_latents: bool = False,253 ):254 sampler = get_sampler_config(params)255 value_dict = {256 "orig_width": image.shape[3] * 8,257 "orig_height": image.shape[2] * 8,258 "target_width": image.shape[3] * 8,259 "target_height": image.shape[2] * 8,260 "prompt": prompt,261 "negative_prompt": negative_prompt,262 "crop_coords_top": 0,263 "crop_coords_left": 0,264 "aesthetic_score": 6.0,265 "negative_aesthetic_score": 2.5,266 }267 268 return do_img2img(269 image,270 self.model,271 sampler,272 value_dict,273 samples,274 skip_encode=True,275 return_latents=return_latents,276 filter=None,277 )278 279 280def get_guider_config(params: SamplingParams):281 if params.guider == Guider.IDENTITY:282 guider_config = {283 "target": "sgm.modules.diffusionmodules.guiders.IdentityGuider"284 }285 elif params.guider == Guider.VANILLA:286 scale = params.scale287 288 thresholder = params.thresholder289 290 if thresholder == Thresholder.NONE:291 dyn_thresh_config = {292 "target": "sgm.modules.diffusionmodules.sampling_utils.NoDynamicThresholding"293 }294 else:295 raise NotImplementedError296 297 guider_config = {298 "target": "sgm.modules.diffusionmodules.guiders.VanillaCFG",299 "params": {"scale": scale, "dyn_thresh_config": dyn_thresh_config},300 }301 else:302 raise NotImplementedError303 return guider_config304 305 306def get_discretization_config(params: SamplingParams):307 if params.discretization == Discretization.LEGACY_DDPM:308 discretization_config = {309 "target": "sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization",310 }311 elif params.discretization == Discretization.EDM:312 discretization_config = {313 "target": "sgm.modules.diffusionmodules.discretizer.EDMDiscretization",314 "params": {315 "sigma_min": params.sigma_min,316 "sigma_max": params.sigma_max,317 "rho": params.rho,318 },319 }320 else:321 raise ValueError(f"unknown discretization {params.discretization}")322 return discretization_config323 324 325def get_sampler_config(params: SamplingParams):326 discretization_config = get_discretization_config(params)327 guider_config = get_guider_config(params)328 sampler = None329 if params.sampler == Sampler.EULER_EDM:330 return EulerEDMSampler(331 num_steps=params.steps,332 discretization_config=discretization_config,333 guider_config=guider_config,334 s_churn=params.s_churn,335 s_tmin=params.s_tmin,336 s_tmax=params.s_tmax,337 s_noise=params.s_noise,338 verbose=True,339 )340 if params.sampler == Sampler.HEUN_EDM:341 return HeunEDMSampler(342 num_steps=params.steps,343 discretization_config=discretization_config,344 guider_config=guider_config,345 s_churn=params.s_churn,346 s_tmin=params.s_tmin,347 s_tmax=params.s_tmax,348 s_noise=params.s_noise,349 verbose=True,350 )351 if params.sampler == Sampler.EULER_ANCESTRAL:352 return EulerAncestralSampler(353 num_steps=params.steps,354 discretization_config=discretization_config,355 guider_config=guider_config,356 eta=params.eta,357 s_noise=params.s_noise,358 verbose=True,359 )360 if params.sampler == Sampler.DPMPP2S_ANCESTRAL:361 return DPMPP2SAncestralSampler(362 num_steps=params.steps,363 discretization_config=discretization_config,364 guider_config=guider_config,365 eta=params.eta,366 s_noise=params.s_noise,367 verbose=True,368 )369 if params.sampler == Sampler.DPMPP2M:370 return DPMPP2MSampler(371 num_steps=params.steps,372 discretization_config=discretization_config,373 guider_config=guider_config,374 verbose=True,375 )376 if params.sampler == Sampler.LINEAR_MULTISTEP:377 return LinearMultistepSampler(378 num_steps=params.steps,379 discretization_config=discretization_config,380 guider_config=guider_config,381 order=params.order,382 verbose=True,383 )384 385 raise ValueError(f"unknown sampler {params.sampler}!")386 