Team Ai
Apppublic

ysharma/ControlNet_Image_Comparison

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
model.py852 linesDownload Raw Back to root
1# This file is adapted from gradio_*.py in https://github.com/lllyasviel/ControlNet/tree/f4748e3630d8141d7765e2bd9b1e348f478477072# The original license file is LICENSE.ControlNet in this repo.3from __future__ import annotations4 5import pathlib6import random7import shlex8import subprocess9import sys10 11import cv212import einops13import numpy as np14import torch15from pytorch_lightning import seed_everything16 17sys.path.append('ControlNet')18 19import config20from annotator.canny import apply_canny21from annotator.hed import apply_hed, nms22from annotator.midas import apply_midas23from annotator.mlsd import apply_mlsd24from annotator.openpose import apply_openpose25from annotator.uniformer import apply_uniformer26from annotator.util import HWC3, resize_image27from cldm.model import create_model, load_state_dict28from ldm.models.diffusion.ddim import DDIMSampler29from share import *30 31from PIL import Image32import gradio as gr33import numpy as np34import base6435 36ORIGINAL_MODEL_NAMES = {37    'canny': 'control_sd15_canny.pth',38    'hough': 'control_sd15_mlsd.pth',39    'hed': 'control_sd15_hed.pth',40    'scribble': 'control_sd15_scribble.pth',41    'pose': 'control_sd15_openpose.pth',42    'seg': 'control_sd15_seg.pth',43    'depth': 'control_sd15_depth.pth',44    'normal': 'control_sd15_normal.pth',45}46ORIGINAL_WEIGHT_ROOT = 'https://huggingface.co/lllyasviel/ControlNet/resolve/main/models/'47 48LIGHTWEIGHT_MODEL_NAMES = {49    'canny': 'control_canny-fp16.safetensors',50    'hough': 'control_mlsd-fp16.safetensors',51    'hed': 'control_hed-fp16.safetensors',52    'scribble': 'control_scribble-fp16.safetensors',53    'pose': 'control_openpose-fp16.safetensors',54    'seg': 'control_seg-fp16.safetensors',55    'depth': 'control_depth-fp16.safetensors',56    'normal': 'control_normal-fp16.safetensors',57}58LIGHTWEIGHT_WEIGHT_ROOT = 'https://huggingface.co/webui/ControlNet-modules-safetensors/resolve/main/'59 60 61class Model:62    def __init__(self,63                 model_config_path: str = 'ControlNet/models/cldm_v15.yaml',64                 model_dir: str = 'models',65                 use_lightweight: bool = True):66        self.device = torch.device(67            'cuda:0' if torch.cuda.is_available() else 'cpu')68        self.model = create_model(model_config_path).to(self.device)69        self.ddim_sampler = DDIMSampler(self.model)70        self.task_name = ''71 72        self.model_dir = pathlib.Path(model_dir)73        self.model_dir.mkdir(exist_ok=True, parents=True)74 75        self.use_lightweight = use_lightweight76        if use_lightweight:77            self.model_names = LIGHTWEIGHT_MODEL_NAMES78            self.weight_root = LIGHTWEIGHT_WEIGHT_ROOT79            base_model_url = 'https://huggingface.co/runwayml/stable-diffusion-v1-5/resolve/main/v1-5-pruned-emaonly.safetensors'80            self.load_base_model(base_model_url)81        else:82            self.model_names = ORIGINAL_MODEL_NAMES83            self.weight_root = ORIGINAL_WEIGHT_ROOT84 85        self.download_models()86 87    def download_base_model(self, model_url: str) -> pathlib.Path:88        model_name = model_url.split('/')[-1]89        out_path = self.model_dir / model_name90        if not out_path.exists():91            subprocess.run(shlex.split(f'wget {model_url} -O {out_path}'))92        return out_path93 94    def load_base_model(self, model_url: str) -> None:95        model_path = self.download_base_model(model_url)96        self.model.load_state_dict(load_state_dict(model_path,97                                                   location=self.device.type),98                                   strict=False)99 100    def load_weight(self, task_name: str) -> None:101        if task_name == self.task_name:102            return103        weight_path = self.get_weight_path(task_name)104        if not self.use_lightweight:105            self.model.load_state_dict(106                load_state_dict(weight_path, location=self.device))107        else:108            self.model.control_model.load_state_dict(109                load_state_dict(weight_path, location=self.device.type))110        self.task_name = task_name111 112    def get_weight_path(self, task_name: str) -> str:113        if 'scribble' in task_name:114            task_name = 'scribble'115        return f'{self.model_dir}/{self.model_names[task_name]}'116 117    def download_models(self) -> None:118        self.model_dir.mkdir(exist_ok=True, parents=True)119        for name in self.model_names.values():120            out_path = self.model_dir / name121            if out_path.exists():122                continue123            subprocess.run(124                shlex.split(f'wget {self.weight_root}{name} -O {out_path}'))125 126    @torch.inference_mode()127    def process_canny(self, input_image, prompt, a_prompt, n_prompt,128                      num_samples, image_resolution, ddim_steps, scale, seed,129                      eta, low_threshold, high_threshold):130        self.load_weight('canny')131 132        img = resize_image(HWC3(input_image), image_resolution)133        H, W, C = img.shape134 135        detected_map = apply_canny(img, low_threshold, high_threshold)136        detected_map = HWC3(detected_map)137 138        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0139        control = torch.stack([control for _ in range(num_samples)], dim=0)140        control = einops.rearrange(control, 'b h w c -> b c h w').clone()141 142        if seed == -1:143            seed = random.randint(0, 65535)144        seed_everything(seed)145 146        if config.save_memory:147            self.model.low_vram_shift(is_diffusing=False)148 149        cond = {150            'c_concat': [control],151            'c_crossattn': [152                self.model.get_learned_conditioning(153                    [prompt + ', ' + a_prompt] * num_samples)154            ]155        }156        un_cond = {157            'c_concat': [control],158            'c_crossattn':159            [self.model.get_learned_conditioning([n_prompt] * num_samples)]160        }161        shape = (4, H // 8, W // 8)162 163        if config.save_memory:164            self.model.low_vram_shift(is_diffusing=True)165 166        samples, intermediates = self.ddim_sampler.sample(167            ddim_steps,168            num_samples,169            shape,170            cond,171            verbose=False,172            eta=eta,173            unconditional_guidance_scale=scale,174            unconditional_conditioning=un_cond)175 176        if config.save_memory:177            self.model.low_vram_shift(is_diffusing=False)178 179        x_samples = self.model.decode_first_stage(samples)180        x_samples = (181            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +182            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)183 184        results = [x_samples[i] for i in range(num_samples)]185        return [255 - detected_map] + results186 187    @torch.inference_mode()188    def process_hough(self, input_image, prompt, a_prompt, n_prompt,189                      num_samples, image_resolution, detect_resolution,190                      ddim_steps, scale, seed, eta, value_threshold,191                      distance_threshold):192        self.load_weight('hough')193 194        input_image = HWC3(input_image)195        detected_map = apply_mlsd(resize_image(input_image, detect_resolution),196                                  value_threshold, distance_threshold)197        detected_map = HWC3(detected_map)198        img = resize_image(input_image, image_resolution)199        H, W, C = img.shape200 201        detected_map = cv2.resize(detected_map, (W, H),202                                  interpolation=cv2.INTER_NEAREST)203 204        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0205        control = torch.stack([control for _ in range(num_samples)], dim=0)206        control = einops.rearrange(control, 'b h w c -> b c h w').clone()207 208        if seed == -1:209            seed = random.randint(0, 65535)210        seed_everything(seed)211 212        if config.save_memory:213            self.model.low_vram_shift(is_diffusing=False)214 215        cond = {216            'c_concat': [control],217            'c_crossattn': [218                self.model.get_learned_conditioning(219                    [prompt + ', ' + a_prompt] * num_samples)220            ]221        }222        un_cond = {223            'c_concat': [control],224            'c_crossattn':225            [self.model.get_learned_conditioning([n_prompt] * num_samples)]226        }227        shape = (4, H // 8, W // 8)228 229        if config.save_memory:230            self.model.low_vram_shift(is_diffusing=True)231 232        samples, intermediates = self.ddim_sampler.sample(233            ddim_steps,234            num_samples,235            shape,236            cond,237            verbose=False,238            eta=eta,239            unconditional_guidance_scale=scale,240            unconditional_conditioning=un_cond)241 242        if config.save_memory:243            self.model.low_vram_shift(is_diffusing=False)244 245        x_samples = self.model.decode_first_stage(samples)246        x_samples = (247            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +248            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)249 250        results = [x_samples[i] for i in range(num_samples)]251        return [252            255 - cv2.dilate(detected_map,253                             np.ones(shape=(3, 3), dtype=np.uint8),254                             iterations=1)255        ] + results256 257    @torch.inference_mode()258    def process_hed(self, input_image, prompt, a_prompt, n_prompt, num_samples,259                    image_resolution, detect_resolution, ddim_steps, scale,260                    seed, eta):261        self.load_weight('hed')262 263        input_image = HWC3(input_image)264        detected_map = apply_hed(resize_image(input_image, detect_resolution))265        detected_map = HWC3(detected_map)266        img = resize_image(input_image, image_resolution)267        H, W, C = img.shape268 269        detected_map = cv2.resize(detected_map, (W, H),270                                  interpolation=cv2.INTER_LINEAR)271 272        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0273        control = torch.stack([control for _ in range(num_samples)], dim=0)274        control = einops.rearrange(control, 'b h w c -> b c h w').clone()275 276        if seed == -1:277            seed = random.randint(0, 65535)278        seed_everything(seed)279 280        if config.save_memory:281            self.model.low_vram_shift(is_diffusing=False)282 283        cond = {284            'c_concat': [control],285            'c_crossattn': [286                self.model.get_learned_conditioning(287                    [prompt + ', ' + a_prompt] * num_samples)288            ]289        }290        un_cond = {291            'c_concat': [control],292            'c_crossattn':293            [self.model.get_learned_conditioning([n_prompt] * num_samples)]294        }295        shape = (4, H // 8, W // 8)296 297        if config.save_memory:298            self.model.low_vram_shift(is_diffusing=True)299 300        samples, intermediates = self.ddim_sampler.sample(301            ddim_steps,302            num_samples,303            shape,304            cond,305            verbose=False,306            eta=eta,307            unconditional_guidance_scale=scale,308            unconditional_conditioning=un_cond)309 310        if config.save_memory:311            self.model.low_vram_shift(is_diffusing=False)312 313        x_samples = self.model.decode_first_stage(samples)314        x_samples = (315            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +316            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)317 318        results = [x_samples[i] for i in range(num_samples)]319        return [detected_map] + results320 321    @torch.inference_mode()322    def process_scribble(self, input_image, prompt, a_prompt, n_prompt,323                         num_samples, image_resolution, ddim_steps, scale,324                         seed, eta):325        self.load_weight('scribble')326 327        img = resize_image(HWC3(input_image), image_resolution)328        H, W, C = img.shape329 330        detected_map = np.zeros_like(img, dtype=np.uint8)331        detected_map[np.min(img, axis=2) < 127] = 255332 333        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0334        control = torch.stack([control for _ in range(num_samples)], dim=0)335        control = einops.rearrange(control, 'b h w c -> b c h w').clone()336 337        if seed == -1:338            seed = random.randint(0, 65535)339        seed_everything(seed)340 341        if config.save_memory:342            self.model.low_vram_shift(is_diffusing=False)343 344        cond = {345            'c_concat': [control],346            'c_crossattn': [347                self.model.get_learned_conditioning(348                    [prompt + ', ' + a_prompt] * num_samples)349            ]350        }351        un_cond = {352            'c_concat': [control],353            'c_crossattn':354            [self.model.get_learned_conditioning([n_prompt] * num_samples)]355        }356        shape = (4, H // 8, W // 8)357 358        if config.save_memory:359            self.model.low_vram_shift(is_diffusing=True)360 361        samples, intermediates = self.ddim_sampler.sample(362            ddim_steps,363            num_samples,364            shape,365            cond,366            verbose=False,367            eta=eta,368            unconditional_guidance_scale=scale,369            unconditional_conditioning=un_cond)370 371        if config.save_memory:372            self.model.low_vram_shift(is_diffusing=False)373 374        x_samples = self.model.decode_first_stage(samples)375        x_samples = (376            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +377            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)378 379        results = [x_samples[i] for i in range(num_samples)]380        return [255 - detected_map] + results381 382    @torch.inference_mode()383    def process_scribble_interactive(self, input_image, prompt, a_prompt,384                                     n_prompt, num_samples, image_resolution,385                                     ddim_steps, scale, seed, eta):386        self.load_weight('scribble')387 388        img = resize_image(HWC3(input_image['mask'][:, :, 0]),389                           image_resolution)390        H, W, C = img.shape391 392        detected_map = np.zeros_like(img, dtype=np.uint8)393        detected_map[np.min(img, axis=2) > 127] = 255394 395        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0396        control = torch.stack([control for _ in range(num_samples)], dim=0)397        control = einops.rearrange(control, 'b h w c -> b c h w').clone()398 399        if seed == -1:400            seed = random.randint(0, 65535)401        seed_everything(seed)402 403        if config.save_memory:404            self.model.low_vram_shift(is_diffusing=False)405 406        cond = {407            'c_concat': [control],408            'c_crossattn': [409                self.model.get_learned_conditioning(410                    [prompt + ', ' + a_prompt] * num_samples)411            ]412        }413        un_cond = {414            'c_concat': [control],415            'c_crossattn':416            [self.model.get_learned_conditioning([n_prompt] * num_samples)]417        }418        shape = (4, H // 8, W // 8)419 420        if config.save_memory:421            self.model.low_vram_shift(is_diffusing=True)422 423        samples, intermediates = self.ddim_sampler.sample(424            ddim_steps,425            num_samples,426            shape,427            cond,428            verbose=False,429            eta=eta,430            unconditional_guidance_scale=scale,431            unconditional_conditioning=un_cond)432 433        if config.save_memory:434            self.model.low_vram_shift(is_diffusing=False)435 436        x_samples = self.model.decode_first_stage(samples)437        x_samples = (438            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +439            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)440 441        results = [x_samples[i] for i in range(num_samples)]442        return [255 - detected_map] + results443 444    @torch.inference_mode()445    def process_fake_scribble(self, input_image, prompt, a_prompt, n_prompt,446                              num_samples, image_resolution, detect_resolution,447                              ddim_steps, scale, seed, eta):448        self.load_weight('scribble')449 450        input_image = HWC3(input_image)451        detected_map = apply_hed(resize_image(input_image, detect_resolution))452        detected_map = HWC3(detected_map)453        img = resize_image(input_image, image_resolution)454        H, W, C = img.shape455 456        detected_map = cv2.resize(detected_map, (W, H),457                                  interpolation=cv2.INTER_LINEAR)458        detected_map = nms(detected_map, 127, 3.0)459        detected_map = cv2.GaussianBlur(detected_map, (0, 0), 3.0)460        detected_map[detected_map > 4] = 255461        detected_map[detected_map < 255] = 0462 463        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0464        control = torch.stack([control for _ in range(num_samples)], dim=0)465        control = einops.rearrange(control, 'b h w c -> b c h w').clone()466 467        if seed == -1:468            seed = random.randint(0, 65535)469        seed_everything(seed)470 471        if config.save_memory:472            self.model.low_vram_shift(is_diffusing=False)473 474        cond = {475            'c_concat': [control],476            'c_crossattn': [477                self.model.get_learned_conditioning(478                    [prompt + ', ' + a_prompt] * num_samples)479            ]480        }481        un_cond = {482            'c_concat': [control],483            'c_crossattn':484            [self.model.get_learned_conditioning([n_prompt] * num_samples)]485        }486        shape = (4, H // 8, W // 8)487 488        if config.save_memory:489            self.model.low_vram_shift(is_diffusing=True)490 491        samples, intermediates = self.ddim_sampler.sample(492            ddim_steps,493            num_samples,494            shape,495            cond,496            verbose=False,497            eta=eta,498            unconditional_guidance_scale=scale,499            unconditional_conditioning=un_cond)500 501        if config.save_memory:502            self.model.low_vram_shift(is_diffusing=False)503 504        x_samples = self.model.decode_first_stage(samples)505        x_samples = (506            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +507            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)508 509        results = [x_samples[i] for i in range(num_samples)]510        return [255 - detected_map] + results511 512    @torch.inference_mode()513    def process_pose(self, input_image, prompt, a_prompt, n_prompt,514                     num_samples, image_resolution, detect_resolution,515                     ddim_steps, scale, seed, eta):516        self.load_weight('pose')517 518        input_image = HWC3(input_image)519        detected_map, _ = apply_openpose(520            resize_image(input_image, detect_resolution))521        detected_map = HWC3(detected_map)522        img = resize_image(input_image, image_resolution)523        H, W, C = img.shape524 525        detected_map = cv2.resize(detected_map, (W, H),526                                  interpolation=cv2.INTER_NEAREST)527 528        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0529        control = torch.stack([control for _ in range(num_samples)], dim=0)530        control = einops.rearrange(control, 'b h w c -> b c h w').clone()531 532        if seed == -1:533            seed = random.randint(0, 65535)534        seed_everything(seed)535 536        if config.save_memory:537            self.model.low_vram_shift(is_diffusing=False)538 539        cond = {540            'c_concat': [control],541            'c_crossattn': [542                self.model.get_learned_conditioning(543                    [prompt + ', ' + a_prompt] * num_samples)544            ]545        }546        un_cond = {547            'c_concat': [control],548            'c_crossattn':549            [self.model.get_learned_conditioning([n_prompt] * num_samples)]550        }551        shape = (4, H // 8, W // 8)552 553        if config.save_memory:554            self.model.low_vram_shift(is_diffusing=True)555 556        samples, intermediates = self.ddim_sampler.sample(557            ddim_steps,558            num_samples,559            shape,560            cond,561            verbose=False,562            eta=eta,563            unconditional_guidance_scale=scale,564            unconditional_conditioning=un_cond)565 566        if config.save_memory:567            self.model.low_vram_shift(is_diffusing=False)568 569        x_samples = self.model.decode_first_stage(samples)570        x_samples = (571            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +572            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)573 574        results = [x_samples[i] for i in range(num_samples)]575        return [detected_map] + results576 577    @torch.inference_mode()578    def process_seg(self, input_image, prompt, a_prompt, n_prompt, num_samples,579                    image_resolution, detect_resolution, ddim_steps, scale,580                    seed, eta):581        self.load_weight('seg')582 583        input_image = HWC3(input_image)584        detected_map = apply_uniformer(585            resize_image(input_image, detect_resolution))586        img = resize_image(input_image, image_resolution)587        H, W, C = img.shape588 589        detected_map = cv2.resize(detected_map, (W, H),590                                  interpolation=cv2.INTER_NEAREST)591 592        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0593        control = torch.stack([control for _ in range(num_samples)], dim=0)594        control = einops.rearrange(control, 'b h w c -> b c h w').clone()595 596        if seed == -1:597            seed = random.randint(0, 65535)598        seed_everything(seed)599 600        if config.save_memory:601            self.model.low_vram_shift(is_diffusing=False)602 603        cond = {604            'c_concat': [control],605            'c_crossattn': [606                self.model.get_learned_conditioning(607                    [prompt + ', ' + a_prompt] * num_samples)608            ]609        }610        un_cond = {611            'c_concat': [control],612            'c_crossattn':613            [self.model.get_learned_conditioning([n_prompt] * num_samples)]614        }615        shape = (4, H // 8, W // 8)616 617        if config.save_memory:618            self.model.low_vram_shift(is_diffusing=True)619 620        samples, intermediates = self.ddim_sampler.sample(621            ddim_steps,622            num_samples,623            shape,624            cond,625            verbose=False,626            eta=eta,627            unconditional_guidance_scale=scale,628            unconditional_conditioning=un_cond)629 630        if config.save_memory:631            self.model.low_vram_shift(is_diffusing=False)632 633        x_samples = self.model.decode_first_stage(samples)634        x_samples = (635            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +636            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)637 638        results = [x_samples[i] for i in range(num_samples)]639 640        tmp = """print(f"type of results ^^ - {type(results)}")641        print(f"length of results list ^^ - {len(results)}")642        print(f"value of results[0] ^^ - {results[0]}")643        filename = results[0] #['name']  644        #def encode(img_array):645        print(f"type of input_image ^^ - {type(input_image)}")646        # Convert NumPy array to image647        img = Image.fromarray(input_image)648        # Save image to file649        img_path = "temp_image.jpeg"650        img.save(img_path)651        # Encode image file using Base64652        with open(img_path, "rb") as image_file:653            encoded_string = base64.b64encode(image_file.read()).decode("utf-8")654        # Print the partial encoded string655        print(encoded_string[:20])656        #return encoded_string657        #def create_imgcomp(input_image, filename):658        #encoded_string = encode(input_image)659        #dummyfun(result_gallery)660        htmltag = '<img src= "data:image/jpeg;base64,' + encoded_string + '" alt="Original Image"/></div> <img src= "https://ysharma-controlnet-image-comparison.hf.space/file=' + filename + '" alt="Control Net Image"/>'661        #https://ysharma-controlnet-image-comparison.hf.space/file=/tmp/tmpg4qx22xy.png - sample662        print(f"htmltag is ^^ - {htmltag}")663        664        desc = 665            <!DOCTYPE html>666            <html lang="en">667            <head>668            	<style>669            		body {670            			background: rgb(17, 17, 17);671            		}672            		673            		.image-slider {674            			margin-left: 3rem;675            			position: relative;676            			display: inline-block;677            			line-height: 0;678            		}679            		680            		.image-slider img {681            			user-select: none;682            			max-width: 400px;683            		}684            		685            		.image-slider > div {686            			position: absolute;687            			width: 25px;688            			max-width: 100%;689            			overflow: hidden;690            			resize: horizontal;691            		}692            		693            		.image-slider > div:before {694            			content: '';695            			display: block;696            			width: 13px;697            			height: 13px;698            			overflow: hidden;699            			position: absolute;700            			resize: horizontal;701            			right: 3px;702            			bottom: 3px;703            			background-clip: content-box;704            			background: linear-gradient(-45deg, black 50%, transparent 0);705            			-webkit-filter: drop-shadow(0 0 2px black);706            			filter: drop-shadow(0 0 2px black);707            		}708            	</style>709            </head>710            <body>711            	<div style="margin: 3rem;712            				font-family: Roboto, sans-serif">713            		<h4 style="color: green"> Observe the Ingenuity of ControlNet by comparing Input and Output images</h4>714            		</div> <div> <div class="image-slider"> <div>  + htmltag + "</div> </div> </body> </html> "715        #return desc716        """717                        718        msg = '<h4 style="color: green"> Observe the Ingenuity of ControlNet by comparing Input and Output images</h4>'             719        return results[0], msg #[detected_map] + results, desc720 721    @torch.inference_mode()722    def process_depth(self, input_image, prompt, a_prompt, n_prompt,723                      num_samples, image_resolution, detect_resolution,724                      ddim_steps, scale, seed, eta):725        self.load_weight('depth')726 727        input_image = HWC3(input_image)728        detected_map, _ = apply_midas(729            resize_image(input_image, detect_resolution))730        detected_map = HWC3(detected_map)731        img = resize_image(input_image, image_resolution)732        H, W, C = img.shape733 734        detected_map = cv2.resize(detected_map, (W, H),735                                  interpolation=cv2.INTER_LINEAR)736 737        control = torch.from_numpy(detected_map.copy()).float().cuda() / 255.0738        control = torch.stack([control for _ in range(num_samples)], dim=0)739        control = einops.rearrange(control, 'b h w c -> b c h w').clone()740 741        if seed == -1:742            seed = random.randint(0, 65535)743        seed_everything(seed)744 745        if config.save_memory:746            self.model.low_vram_shift(is_diffusing=False)747 748        cond = {749            'c_concat': [control],750            'c_crossattn': [751                self.model.get_learned_conditioning(752                    [prompt + ', ' + a_prompt] * num_samples)753            ]754        }755        un_cond = {756            'c_concat': [control],757            'c_crossattn':758            [self.model.get_learned_conditioning([n_prompt] * num_samples)]759        }760        shape = (4, H // 8, W // 8)761 762        if config.save_memory:763            self.model.low_vram_shift(is_diffusing=True)764 765        samples, intermediates = self.ddim_sampler.sample(766            ddim_steps,767            num_samples,768            shape,769            cond,770            verbose=False,771            eta=eta,772            unconditional_guidance_scale=scale,773            unconditional_conditioning=un_cond)774 775        if config.save_memory:776            self.model.low_vram_shift(is_diffusing=False)777 778        x_samples = self.model.decode_first_stage(samples)779        x_samples = (780            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +781            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)782 783        results = [x_samples[i] for i in range(num_samples)]784        return [detected_map] + results785 786    @torch.inference_mode()787    def process_normal(self, input_image, prompt, a_prompt, n_prompt,788                       num_samples, image_resolution, detect_resolution,789                       ddim_steps, scale, seed, eta, bg_threshold):790        self.load_weight('normal')791 792        input_image = HWC3(input_image)793        _, detected_map = apply_midas(resize_image(input_image,794                                                   detect_resolution),795                                      bg_th=bg_threshold)796        detected_map = HWC3(detected_map)797        img = resize_image(input_image, image_resolution)798        H, W, C = img.shape799 800        detected_map = cv2.resize(detected_map, (W, H),801                                  interpolation=cv2.INTER_LINEAR)802 803        control = torch.from_numpy(804            detected_map[:, :, ::-1].copy()).float().cuda() / 255.0805        control = torch.stack([control for _ in range(num_samples)], dim=0)806        control = einops.rearrange(control, 'b h w c -> b c h w').clone()807 808        if seed == -1:809            seed = random.randint(0, 65535)810        seed_everything(seed)811 812        if config.save_memory:813            self.model.low_vram_shift(is_diffusing=False)814 815        cond = {816            'c_concat': [control],817            'c_crossattn': [818                self.model.get_learned_conditioning(819                    [prompt + ', ' + a_prompt] * num_samples)820            ]821        }822        un_cond = {823            'c_concat': [control],824            'c_crossattn':825            [self.model.get_learned_conditioning([n_prompt] * num_samples)]826        }827        shape = (4, H // 8, W // 8)828 829        if config.save_memory:830            self.model.low_vram_shift(is_diffusing=True)831 832        samples, intermediates = self.ddim_sampler.sample(833            ddim_steps,834            num_samples,835            shape,836            cond,837            verbose=False,838            eta=eta,839            unconditional_guidance_scale=scale,840            unconditional_conditioning=un_cond)841 842        if config.save_memory:843            self.model.low_vram_shift(is_diffusing=False)844 845        x_samples = self.model.decode_first_stage(samples)846        x_samples = (847            einops.rearrange(x_samples, 'b c h w -> b h w c') * 127.5 +848            127.5).cpu().numpy().clip(0, 255).astype(np.uint8)849 850        results = [x_samples[i] for i in range(num_samples)]851        return [detected_map] + results852