coding-alt/IF
1
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 