Team Ai
Apppublic

coding-alt/IF

sourceHugging Faceotherupdated 3y agoView on Hugging Face
1likes
model.py314 linesDownload Raw Back to root
1from __future__ import annotations2 3import gc4import json5import tempfile6from typing import Generator7 8import numpy as np9import PIL.Image10import torch11from diffusers import DiffusionPipeline, StableDiffusionUpscalePipeline12from diffusers.pipelines.deepfloyd_if import (fast27_timesteps,13                                              smart27_timesteps,14                                              smart50_timesteps,15                                              smart100_timesteps,16                                              smart185_timesteps)17 18from settings import (DISABLE_AUTOMATIC_CPU_OFFLOAD, DISABLE_SD_X4_UPSCALER,19                      HF_TOKEN, MAX_NUM_IMAGES, MAX_NUM_STEPS, MAX_SEED,20                      RUN_GARBAGE_COLLECTION)21 22 23class Model:24    def __init__(self):25        self.device = torch.device(26            'cuda:0' if torch.cuda.is_available() else 'cpu')27        self.pipe = None28        self.super_res_1_pipe = None29        self.super_res_2_pipe = None30        self.watermark_image = None31 32        if torch.cuda.is_available():33            self.load_weights()34            self.watermark_image = PIL.Image.fromarray(35                self.pipe.watermarker.watermark_image.to(36                    torch.uint8).cpu().numpy(),37                mode='RGBA')38 39    def load_weights(self) -> None:40        self.pipe = DiffusionPipeline.from_pretrained(41            'DeepFloyd/IF-I-XL-v1.0',42            torch_dtype=torch.float16,43            variant='fp16',44            use_safetensors=True,45            use_auth_token=HF_TOKEN)46        self.super_res_1_pipe = DiffusionPipeline.from_pretrained(47            'DeepFloyd/IF-II-L-v1.0',48            text_encoder=None,49            torch_dtype=torch.float16,50            variant='fp16',51            use_safetensors=True,52            use_auth_token=HF_TOKEN)53 54        if not DISABLE_SD_X4_UPSCALER:55            self.super_res_2_pipe = StableDiffusionUpscalePipeline.from_pretrained(56                'stabilityai/stable-diffusion-x4-upscaler',57                torch_dtype=torch.float16)58 59        if DISABLE_AUTOMATIC_CPU_OFFLOAD:60            self.pipe.to(self.device)61            self.super_res_1_pipe.to(self.device)62 63            self.pipe.unet.to(memory_format=torch.channels_last)64            self.pipe.unet = torch.compile(self.pipe.unet, mode="reduce-overhead", fullgraph=True)65 66            if not DISABLE_SD_X4_UPSCALER:67                self.super_res_2_pipe.to(self.device)68        else:69            self.pipe.enable_model_cpu_offload()70            self.super_res_1_pipe.enable_model_cpu_offload()71            if not DISABLE_SD_X4_UPSCALER:72                self.super_res_2_pipe.enable_model_cpu_offload()73 74    def apply_watermark_to_sd_x4_upscaler_results(75            self, images: list[PIL.Image.Image]) -> None:76        w, h = images[0].size77 78        stability_x4_upscaler_sample_size = 12879 80        coef = min(h / stability_x4_upscaler_sample_size,81                   w / stability_x4_upscaler_sample_size)82        img_h, img_w = (int(h / coef), int(w / coef)) if coef < 1 else (h, w)83 84        S1, S2 = 1024**2, img_w * img_h85        K = (S2 / S1)**0.586        watermark_size = int(K * 62)87        watermark_x = img_w - int(14 * K)88        watermark_y = img_h - int(14 * K)89 90        watermark_image = self.watermark_image.copy().resize(91            (watermark_size, watermark_size),92            PIL.Image.Resampling.BICUBIC,93            reducing_gap=None)94 95        for image in images:96            image.paste(watermark_image,97                        box=(98                            watermark_x - watermark_size,99                            watermark_y - watermark_size,100                            watermark_x,101                            watermark_y,102                        ),103                        mask=watermark_image.split()[-1])104 105    @staticmethod106    def to_pil_images(images: torch.Tensor) -> list[PIL.Image.Image]:107        images = (images / 2 + 0.5).clamp(0, 1)108        images = images.cpu().permute(0, 2, 3, 1).float().numpy()109        images = np.round(images * 255).astype(np.uint8)110        return [PIL.Image.fromarray(image) for image in images]111 112    @staticmethod113    def check_seed(seed: int) -> None:114        if not 0 <= seed <= MAX_SEED:115            raise ValueError116 117    @staticmethod118    def check_num_images(num_images: int) -> None:119        if not 1 <= num_images <= MAX_NUM_IMAGES:120            raise ValueError121 122    @staticmethod123    def check_num_inference_steps(num_steps: int) -> None:124        if not 1 <= num_steps <= MAX_NUM_STEPS:125            raise ValueError126 127    @staticmethod128    def get_custom_timesteps(name: str) -> list[int] | None:129        if name == 'none':130            timesteps = None131        elif name == 'fast27':132            timesteps = fast27_timesteps133        elif name == 'smart27':134            timesteps = smart27_timesteps135        elif name == 'smart50':136            timesteps = smart50_timesteps137        elif name == 'smart100':138            timesteps = smart100_timesteps139        elif name == 'smart185':140            timesteps = smart185_timesteps141        else:142            raise ValueError143        return timesteps144 145    @staticmethod146    def run_garbage_collection():147        gc.collect()148        torch.cuda.empty_cache()149 150    def run_stage1(151        self,152        prompt: str,153        negative_prompt: str = '',154        seed: int = 0,155        num_images: int = 1,156        guidance_scale_1: float = 7.0,157        custom_timesteps_1: str = 'smart100',158        num_inference_steps_1: int = 100,159    ) -> tuple[list[PIL.Image.Image], str, str]:160        self.check_seed(seed)161        self.check_num_images(num_images)162        self.check_num_inference_steps(num_inference_steps_1)163 164        if RUN_GARBAGE_COLLECTION:165            self.run_garbage_collection()166 167        generator = torch.Generator(device=self.device).manual_seed(seed)168 169        prompt_embeds, negative_embeds = self.pipe.encode_prompt(170            prompt=prompt, negative_prompt=negative_prompt)171 172        timesteps = self.get_custom_timesteps(custom_timesteps_1)173 174        images = self.pipe(prompt_embeds=prompt_embeds,175                           negative_prompt_embeds=negative_embeds,176                           num_images_per_prompt=num_images,177                           guidance_scale=guidance_scale_1,178                           timesteps=timesteps,179                           num_inference_steps=num_inference_steps_1,180                           generator=generator,181                           output_type='pt').images182        pil_images = self.to_pil_images(images)183        self.pipe.watermarker.apply_watermark(184            pil_images, self.pipe.unet.config.sample_size)185 186        stage1_params = {187            'prompt': prompt,188            'negative_prompt': negative_prompt,189            'seed': seed,190            'num_images': num_images,191            'guidance_scale_1': guidance_scale_1,192            'custom_timesteps_1': custom_timesteps_1,193            'num_inference_steps_1': num_inference_steps_1,194        }195        with tempfile.NamedTemporaryFile(mode='w', delete=False) as param_file:196            param_file.write(json.dumps(stage1_params))197        stage1_result = {198            'prompt_embeds': prompt_embeds,199            'negative_embeds': negative_embeds,200            'images': images,201            'pil_images': pil_images,202        }203        with tempfile.NamedTemporaryFile(delete=False) as result_file:204            torch.save(stage1_result, result_file.name)205        return pil_images, param_file.name, result_file.name206 207    def run_stage2(208        self,209        stage1_result_path: str,210        stage2_index: int,211        seed_2: int = 0,212        guidance_scale_2: float = 4.0,213        custom_timesteps_2: str = 'smart50',214        num_inference_steps_2: int = 50,215        disable_watermark: bool = False,216    ) -> PIL.Image.Image:217        self.check_seed(seed_2)218        self.check_num_inference_steps(num_inference_steps_2)219 220        if RUN_GARBAGE_COLLECTION:221            self.run_garbage_collection()222 223        generator = torch.Generator(device=self.device).manual_seed(seed_2)224 225        stage1_result = torch.load(stage1_result_path)226        prompt_embeds = stage1_result['prompt_embeds']227        negative_embeds = stage1_result['negative_embeds']228        images = stage1_result['images']229        images = images[[stage2_index]]230 231        timesteps = self.get_custom_timesteps(custom_timesteps_2)232 233        out = self.super_res_1_pipe(image=images,234                                    prompt_embeds=prompt_embeds,235                                    negative_prompt_embeds=negative_embeds,236                                    num_images_per_prompt=1,237                                    guidance_scale=guidance_scale_2,238                                    timesteps=timesteps,239                                    num_inference_steps=num_inference_steps_2,240                                    generator=generator,241                                    output_type='pt',242                                    noise_level=250).images243        pil_images = self.to_pil_images(out)244 245        if disable_watermark:246            return pil_images[0]247 248        self.super_res_1_pipe.watermarker.apply_watermark(249            pil_images, self.super_res_1_pipe.unet.config.sample_size)250        return pil_images[0]251 252    def run_stage3(253        self,254        image: PIL.Image.Image,255        prompt: str = '',256        negative_prompt: str = '',257        seed_3: int = 0,258        guidance_scale_3: float = 9.0,259        num_inference_steps_3: int = 75,260    ) -> PIL.Image.Image:261        self.check_seed(seed_3)262        self.check_num_inference_steps(num_inference_steps_3)263 264        if RUN_GARBAGE_COLLECTION:265            self.run_garbage_collection()266 267        generator = torch.Generator(device=self.device).manual_seed(seed_3)268        out = self.super_res_2_pipe(image=image,269                                    prompt=prompt,270                                    negative_prompt=negative_prompt,271                                    num_images_per_prompt=1,272                                    guidance_scale=guidance_scale_3,273                                    num_inference_steps=num_inference_steps_3,274                                    generator=generator,275                                    noise_level=100).images276        self.apply_watermark_to_sd_x4_upscaler_results(out)277        return out[0]278 279    def run_stage2_3(280        self,281        stage1_result_path: str,282        stage2_index: int,283        seed_2: int = 0,284        guidance_scale_2: float = 4.0,285        custom_timesteps_2: str = 'smart50',286        num_inference_steps_2: int = 50,287        prompt: str = '',288        negative_prompt: str = '',289        seed_3: int = 0,290        guidance_scale_3: float = 9.0,291        num_inference_steps_3: int = 75,292    ) -> Generator[PIL.Image.Image]:293        self.check_seed(seed_3)294        self.check_num_inference_steps(num_inference_steps_3)295 296        out_image = self.run_stage2(297            stage1_result_path=stage1_result_path,298            stage2_index=stage2_index,299            seed_2=seed_2,300            guidance_scale_2=guidance_scale_2,301            custom_timesteps_2=custom_timesteps_2,302            num_inference_steps_2=num_inference_steps_2,303            disable_watermark=True)304        temp_image = out_image.copy()305        self.super_res_1_pipe.watermarker.apply_watermark(306            [temp_image], self.super_res_1_pipe.unet.config.sample_size)307        yield temp_image308        yield self.run_stage3(image=out_image,309                              prompt=prompt,310                              negative_prompt=negative_prompt,311                              seed_3=seed_3,312                              guidance_scale_3=guidance_scale_3,313                              num_inference_steps_3=num_inference_steps_3)314