CoolFace
Apppublic

zparadox/stable-video-diffusion

sourceHugging Faceotherupdated 3y agoView on Hugging Face
0likes
api.py386 linesDownload Raw Back to inference
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