ysharma/ControlNet_Image_Comparison
3
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 