Team Ai
Apppublic

hysts/diffusers-anime-faces

sourceHugging Faceupdated 4y agoView on Hugging Face
19likes
model.py188 linesDownload Raw Back to root
1from __future__ import annotations2 3import logging4import os5import random6import sys7import tempfile8 9import gradio as gr10import imageio11import numpy as np12import PIL.Image13import torch14import tqdm.auto15from diffusers import (DDIMPipeline, DDIMScheduler, DDPMPipeline,16                       DiffusionPipeline, PNDMPipeline, PNDMScheduler)17 18HF_TOKEN = os.environ['HF_TOKEN']19 20formatter = logging.Formatter(21    '[%(asctime)s] %(name)s %(levelname)s: %(message)s',22    datefmt='%Y-%m-%d %H:%M:%S')23stream_handler = logging.StreamHandler(stream=sys.stdout)24stream_handler.setLevel(logging.INFO)25stream_handler.setFormatter(formatter)26logger = logging.getLogger(__name__)27logger.setLevel(logging.INFO)28logger.propagate = False29logger.addHandler(stream_handler)30 31 32class Model:33 34    MODEL_NAMES = [35        'ddpm-128-exp000',36    ]37 38    def __init__(self, device: str | torch.device):39        self.device = torch.device(device)40        self._download_all_models()41 42        self.model_name = self.MODEL_NAMES[0]43        self.scheduler_type = 'DDIM'44        self.pipeline = self._load_pipeline(self.model_name,45                                            self.scheduler_type)46        self.rng = random.Random()47 48        self.real_esrgan = gr.Interface.load('spaces/hysts/Real-ESRGAN-anime')49 50    @staticmethod51    def _load_pipeline(model_name: str,52                       scheduler_type: str) -> DiffusionPipeline:53        repo_id = f'hysts/diffusers-anime-faces-{model_name}'54        if scheduler_type == 'DDPM':55            pipeline = DDPMPipeline.from_pretrained(repo_id,56                                                    use_auth_token=HF_TOKEN)57        elif scheduler_type == 'DDIM':58            pipeline = DDIMPipeline.from_pretrained(repo_id,59                                                    use_auth_token=HF_TOKEN)60            pipeline.scheduler = DDIMScheduler.from_config(61                repo_id, subfolder='scheduler', use_auth_token=HF_TOKEN)62        elif scheduler_type == 'PNDM':63            pipeline = PNDMPipeline.from_pretrained(repo_id,64                                                    use_auth_token=HF_TOKEN)65            pipeline.scheduler = PNDMScheduler.from_config(66                repo_id, subfolder='scheduler', use_auth_token=HF_TOKEN)67        else:68            raise ValueError69        return pipeline70 71    def set_pipeline(self, model_name: str, scheduler_type: str) -> None:72        logger.info('--- set_pipeline ---')73        logger.info(f'{model_name=}, {scheduler_type=}')74 75        if model_name == self.model_name and scheduler_type == self.scheduler_type:76            logger.info('Skipping')77            logger.info('--- done ---')78            return79        self.model_name = model_name80        self.scheduler_type = scheduler_type81        self.pipeline = self._load_pipeline(model_name, scheduler_type)82 83        logger.info('--- done ---')84 85    def _download_all_models(self) -> None:86        for name in self.MODEL_NAMES:87            self._load_pipeline(name, 'DDPM')88 89    def generate(self,90                 seed: int,91                 num_steps: int,92                 num_images: int = 1) -> list[PIL.Image.Image]:93        logger.info('--- generate ---')94        logger.info(f'{seed=}, {num_steps=}')95 96        torch.manual_seed(seed)97        if self.scheduler_type == 'DDPM':98            res = self.pipeline(batch_size=num_images,99                                torch_device=self.device)['sample']100        elif self.scheduler_type in ['DDIM', 'PNDM']:101            res = self.pipeline(batch_size=num_images,102                                torch_device=self.device,103                                num_inference_steps=num_steps)['sample']104        else:105            raise ValueError106 107        logger.info('--- done ---')108        return res109 110    @staticmethod111    def postprocess(sample: torch.Tensor) -> np.ndarray:112        res = (sample / 2 + 0.5).clamp(0, 1)113        res = (res * 255).to(torch.uint8)114        res = res.cpu().permute(0, 2, 3, 1).numpy()115        return res116 117    @torch.inference_mode()118    def generate_with_video(self, seed: int,119                            num_steps: int) -> tuple[PIL.Image.Image, str]:120        logger.info('--- generate_with_video ---')121        if self.scheduler_type == 'DDPM':122            num_steps = 1000123            fps = 100124        else:125            fps = 10126        logger.info(f'{seed=}, {num_steps=}')127 128        model = self.pipeline.unet.to(self.device)129        scheduler = self.pipeline.scheduler130        scheduler.set_timesteps(num_inference_steps=num_steps)131        input_shape = (1, model.config.in_channels, model.config.sample_size,132                       model.config.sample_size)133        torch.manual_seed(seed)134 135        out_file = tempfile.NamedTemporaryFile(suffix='.mp4', delete=False)136        writer = imageio.get_writer(out_file.name, fps=fps)137        sample = torch.randn(input_shape).to(self.device)138        for t in tqdm.auto.tqdm(scheduler.timesteps):139            out = model(sample, t)['sample']140            sample = scheduler.step(out, t, sample)['prev_sample']141            res = self.postprocess(sample)[0]142            writer.append_data(res)143        writer.close()144 145        logger.info('--- done ---')146        return PIL.Image.fromarray(res), out_file.name147 148    def superresolve(self, image: PIL.Image.Image) -> PIL.Image.Image:149        logger.info('--- superresolve ---')150 151        with tempfile.NamedTemporaryFile(suffix='.png') as f:152            image.save(f.name)153            out_file = self.real_esrgan(f.name)154 155        logger.info('--- done ---')156        return PIL.Image.open(out_file)157 158    def run(self, model_name: str, scheduler_type: str, num_steps: int,159            randomize_seed: bool,160            seed: int) -> tuple[PIL.Image.Image, PIL.Image.Image, int, str]:161        self.set_pipeline(model_name, scheduler_type)162        if scheduler_type == 'PNDM':163            num_steps = max(4, min(num_steps, 100))164        if randomize_seed:165            seed = self.rng.randint(0, 100000)166        res, filename = self.generate_with_video(seed, num_steps)167        superresolved = self.superresolve(res)168        return superresolved, res, seed, filename169 170    @staticmethod171    def to_grid(images: list[PIL.Image.Image],172                ncols: int = 2) -> PIL.Image.Image:173        images = [np.asarray(image) for image in images]174        nrows = (len(images) + ncols - 1) // ncols175        h, w = images[0].shape[:2]176        if (d := nrows * ncols - len(images)) > 0:177            images += [np.full((h, w, 3), 255, dtype=np.uint8)] * d178        grid = np.asarray(images).reshape(nrows, ncols, h, w, 3).transpose(179            0, 2, 1, 3, 4).reshape(nrows * h, ncols * w, 3)180        return PIL.Image.fromarray(grid)181 182    def run_simple(self) -> tuple[PIL.Image.Image, PIL.Image.Image]:183        self.set_pipeline(self.MODEL_NAMES[0], 'PNDM')184        seed = self.rng.randint(0, np.iinfo(np.uint32).max + 1)185        images = self.generate(seed, num_steps=10, num_images=4)186        superresolved = [self.superresolve(image) for image in images]187        return self.to_grid(superresolved, 2), self.to_grid(images, 2)188